[diffusion] chore: centralize entrypoint API hygiene (#33845)
This commit is contained in:
@@ -1,4 +1 @@
|
||||
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import globally_suppress_loggers
|
||||
|
||||
globally_suppress_loggers()
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
@@ -43,6 +43,7 @@ from sglang.multimodal_gen.runtime.server_warmup import (
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||
GREEN,
|
||||
RESET,
|
||||
globally_suppress_loggers,
|
||||
init_logger,
|
||||
log_batch_completion,
|
||||
log_generation_timer,
|
||||
@@ -125,12 +126,13 @@ class DiffGenerator:
|
||||
Returns:
|
||||
The created DiffGenerator
|
||||
"""
|
||||
globally_suppress_loggers()
|
||||
instance = cls(
|
||||
server_args=server_args,
|
||||
)
|
||||
init_diffusion_tracing(server_args, "DiffGenerator")
|
||||
|
||||
logger.info(f"Local mode: {local_mode}")
|
||||
logger.info("Local mode: %s", local_mode)
|
||||
if local_mode:
|
||||
instance.local_scheduler_process = instance._start_local_server_if_needed()
|
||||
instance.owns_scheduler_client = True
|
||||
|
||||
@@ -38,7 +38,10 @@ from sglang.multimodal_gen.runtime.server_warmup import (
|
||||
run_async_client_warmup,
|
||||
should_run_synthetic_server_warmup,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||
globally_suppress_loggers,
|
||||
init_logger,
|
||||
)
|
||||
from sglang.srt.utils.json_response import orjson_response
|
||||
from sglang.version import __version__
|
||||
|
||||
@@ -374,6 +377,7 @@ def create_app(server_args: ServerArgs):
|
||||
"""
|
||||
Create and configure the FastAPI application instance.
|
||||
"""
|
||||
globally_suppress_loggers()
|
||||
app = FastAPI(lifespan=lifespan)
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
|
||||
@@ -14,7 +14,7 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
|
||||
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
|
||||
from sglang.multimodal_gen.runtime.server_args import get_global_server_args
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.srt.utils.json_response import orjson_response
|
||||
|
||||
@@ -45,6 +45,26 @@ class DiffusionModelCard(ModelCard):
|
||||
pipeline_class: Optional[str] = None
|
||||
|
||||
|
||||
def _build_model_card(server_args: ServerArgs, model_id: str) -> DiffusionModelCard:
|
||||
model_info = get_model_info(
|
||||
server_args.model_path,
|
||||
backend=server_args.backend,
|
||||
model_id=server_args.model_id,
|
||||
)
|
||||
card_kwargs: dict[str, Any] = {
|
||||
"id": model_id,
|
||||
"root": model_id,
|
||||
"num_gpus": server_args.num_gpus,
|
||||
"task_type": server_args.pipeline_config.task_type.name,
|
||||
"dit_precision": server_args.pipeline_config.dit_precision,
|
||||
"vae_precision": server_args.pipeline_config.vae_precision,
|
||||
}
|
||||
if model_info:
|
||||
card_kwargs["pipeline_name"] = model_info.pipeline_cls.pipeline_name
|
||||
card_kwargs["pipeline_class"] = model_info.pipeline_cls.__name__
|
||||
return DiffusionModelCard(**card_kwargs)
|
||||
|
||||
|
||||
async def _handle_lora_request(req: Any, success_msg: str, failure_msg: str):
|
||||
try:
|
||||
output: OutputBatch = await async_scheduler_client.forward(req)
|
||||
@@ -183,27 +203,7 @@ async def available_models():
|
||||
if not server_args:
|
||||
raise HTTPException(status_code=500, detail="Server args not initialized")
|
||||
|
||||
model_info = get_model_info(
|
||||
server_args.model_path,
|
||||
backend=server_args.backend,
|
||||
model_id=server_args.model_id,
|
||||
)
|
||||
|
||||
card_kwargs = {
|
||||
"id": server_args.model_path,
|
||||
"root": server_args.model_path,
|
||||
# Extended diffusion-specific fields
|
||||
"num_gpus": server_args.num_gpus,
|
||||
"task_type": server_args.pipeline_config.task_type.name,
|
||||
"dit_precision": server_args.pipeline_config.dit_precision,
|
||||
"vae_precision": server_args.pipeline_config.vae_precision,
|
||||
}
|
||||
|
||||
if model_info:
|
||||
card_kwargs["pipeline_name"] = model_info.pipeline_cls.pipeline_name
|
||||
card_kwargs["pipeline_class"] = model_info.pipeline_cls.__name__
|
||||
|
||||
model_card = DiffusionModelCard(**card_kwargs)
|
||||
model_card = _build_model_card(server_args, server_args.model_path)
|
||||
|
||||
# Return dict directly to preserve extended fields (ModelList strips them)
|
||||
return {"object": "list", "data": [model_card.model_dump()]}
|
||||
@@ -229,24 +229,5 @@ async def retrieve_model(model: str):
|
||||
status_code=404,
|
||||
)
|
||||
|
||||
model_info = get_model_info(
|
||||
server_args.model_path,
|
||||
backend=server_args.backend,
|
||||
model_id=server_args.model_id,
|
||||
)
|
||||
|
||||
card_kwargs = {
|
||||
"id": model,
|
||||
"root": model,
|
||||
"num_gpus": server_args.num_gpus,
|
||||
"task_type": server_args.pipeline_config.task_type.name,
|
||||
"dit_precision": server_args.pipeline_config.dit_precision,
|
||||
"vae_precision": server_args.pipeline_config.vae_precision,
|
||||
}
|
||||
|
||||
if model_info:
|
||||
card_kwargs["pipeline_name"] = model_info.pipeline_cls.pipeline_name
|
||||
card_kwargs["pipeline_class"] = model_info.pipeline_cls.__name__
|
||||
|
||||
# Return dict to preserve extended fields
|
||||
return DiffusionModelCard(**card_kwargs).model_dump()
|
||||
return _build_model_card(server_args, model).model_dump()
|
||||
|
||||
@@ -119,7 +119,7 @@ async def _dispatch_job_async(job_id: str, batch: Req) -> None:
|
||||
)
|
||||
await MESH_STORE.update_fields(job_id, update_fields)
|
||||
except Exception as e:
|
||||
logger.error(f"{e}")
|
||||
logger.exception("Mesh job %s failed", job_id)
|
||||
await MESH_STORE.update_fields(
|
||||
job_id, {"status": "failed", "error": {"message": str(e)}}
|
||||
)
|
||||
@@ -229,24 +229,8 @@ async def list_meshes(
|
||||
limit: Optional[int] = Query(None, ge=1, le=100),
|
||||
order: Optional[str] = Query("desc"),
|
||||
):
|
||||
order = (order or "desc").lower()
|
||||
if order not in ("asc", "desc"):
|
||||
order = "desc"
|
||||
jobs = await MESH_STORE.list_values()
|
||||
|
||||
reverse = order != "asc"
|
||||
jobs.sort(key=lambda j: j.get("created_at", 0), reverse=reverse)
|
||||
|
||||
if after is not None:
|
||||
try:
|
||||
idx = next(i for i, j in enumerate(jobs) if j["id"] == after)
|
||||
jobs = jobs[idx + 1 :]
|
||||
except StopIteration:
|
||||
jobs = []
|
||||
|
||||
if limit is not None:
|
||||
jobs = jobs[:limit]
|
||||
items = [MeshResponse(**j) for j in jobs]
|
||||
jobs = await MESH_STORE.list_page(after=after, limit=limit, order=order)
|
||||
items = [MeshResponse(**job) for job in jobs]
|
||||
return MeshListResponse(data=items)
|
||||
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import asyncio
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
|
||||
class AsyncDictStore:
|
||||
@@ -36,9 +36,31 @@ class AsyncDictStore:
|
||||
async with self._lock:
|
||||
return self._items.pop(key, None)
|
||||
|
||||
async def list_values(self) -> List[Dict[str, Any]]:
|
||||
async def list_page(
|
||||
self,
|
||||
*,
|
||||
after: Optional[str] = None,
|
||||
limit: Optional[int] = None,
|
||||
order: Optional[str] = "desc",
|
||||
) -> list[Dict[str, Any]]:
|
||||
"""Return a created-time ordered page, with OpenAI-style cursor semantics."""
|
||||
normalized_order = (order or "desc").lower()
|
||||
reverse = normalized_order != "asc"
|
||||
async with self._lock:
|
||||
return list(self._items.values())
|
||||
items = sorted(
|
||||
self._items.values(),
|
||||
key=lambda item: item.get("created_at", 0),
|
||||
reverse=reverse,
|
||||
)
|
||||
|
||||
if after is not None:
|
||||
try:
|
||||
index = next(i for i, item in enumerate(items) if item["id"] == after)
|
||||
items = items[index + 1 :]
|
||||
except StopIteration:
|
||||
return []
|
||||
|
||||
return items[:limit] if limit is not None else items
|
||||
|
||||
|
||||
# Global stores shared by OpenAI entrypoints
|
||||
|
||||
@@ -489,7 +489,7 @@ async def _dispatch_job_async(
|
||||
update_fields.update(final_media_fields)
|
||||
await VIDEO_STORE.update_fields(job_id, update_fields)
|
||||
except Exception as e:
|
||||
logger.error(f"{e}")
|
||||
logger.exception("Video job %s failed", job_id)
|
||||
await VIDEO_STORE.update_fields(
|
||||
job_id,
|
||||
{
|
||||
@@ -842,25 +842,8 @@ async def list_videos(
|
||||
limit: Optional[int] = Query(None, ge=1, le=100),
|
||||
order: Optional[str] = Query("desc"),
|
||||
):
|
||||
# Normalize order
|
||||
order = (order or "desc").lower()
|
||||
if order not in ("asc", "desc"):
|
||||
order = "desc"
|
||||
jobs = await VIDEO_STORE.list_values()
|
||||
|
||||
reverse = order != "asc"
|
||||
jobs.sort(key=lambda j: j.get("created_at", 0), reverse=reverse)
|
||||
|
||||
if after is not None:
|
||||
try:
|
||||
idx = next(i for i, j in enumerate(jobs) if j["id"] == after)
|
||||
jobs = jobs[idx + 1 :]
|
||||
except StopIteration:
|
||||
jobs = []
|
||||
|
||||
if limit is not None:
|
||||
jobs = jobs[:limit]
|
||||
items = [VideoResponse(**j) for j in jobs]
|
||||
jobs = await VIDEO_STORE.list_page(after=after, limit=limit, order=order)
|
||||
items = [VideoResponse(**job) for job in jobs]
|
||||
return VideoListResponse(data=items)
|
||||
|
||||
|
||||
|
||||
@@ -66,7 +66,7 @@ async def run_zeromq_broker(server_args: ServerArgs):
|
||||
# 1. Receive a request from an offline client
|
||||
payload = await socket.recv()
|
||||
request_batch = pickle.loads(payload)
|
||||
logger.info("Broker received an offline job from a client.")
|
||||
logger.debug("Broker received an offline job from a client.")
|
||||
|
||||
# 2. Forward the request to the main Scheduler via the shared client
|
||||
response_batch = await async_scheduler_client.forward(request_batch)
|
||||
|
||||
Reference in New Issue
Block a user