[diffusion] chore: centralize entrypoint API hygiene (#33845)

This commit is contained in:
Mick
2026-08-06 21:53:40 +08:00
committed by GitHub
parent 2132cdef16
commit 183bd80add
8 changed files with 64 additions and 91 deletions
@@ -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)