[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
|
# SPDX-License-Identifier: Apache-2.0
|
||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import globally_suppress_loggers
|
|
||||||
|
|
||||||
globally_suppress_loggers()
|
|
||||||
|
|||||||
@@ -43,6 +43,7 @@ from sglang.multimodal_gen.runtime.server_warmup import (
|
|||||||
from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
from sglang.multimodal_gen.runtime.utils.logging_utils import (
|
||||||
GREEN,
|
GREEN,
|
||||||
RESET,
|
RESET,
|
||||||
|
globally_suppress_loggers,
|
||||||
init_logger,
|
init_logger,
|
||||||
log_batch_completion,
|
log_batch_completion,
|
||||||
log_generation_timer,
|
log_generation_timer,
|
||||||
@@ -125,12 +126,13 @@ class DiffGenerator:
|
|||||||
Returns:
|
Returns:
|
||||||
The created DiffGenerator
|
The created DiffGenerator
|
||||||
"""
|
"""
|
||||||
|
globally_suppress_loggers()
|
||||||
instance = cls(
|
instance = cls(
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
)
|
)
|
||||||
init_diffusion_tracing(server_args, "DiffGenerator")
|
init_diffusion_tracing(server_args, "DiffGenerator")
|
||||||
|
|
||||||
logger.info(f"Local mode: {local_mode}")
|
logger.info("Local mode: %s", local_mode)
|
||||||
if local_mode:
|
if local_mode:
|
||||||
instance.local_scheduler_process = instance._start_local_server_if_needed()
|
instance.local_scheduler_process = instance._start_local_server_if_needed()
|
||||||
instance.owns_scheduler_client = True
|
instance.owns_scheduler_client = True
|
||||||
|
|||||||
@@ -38,7 +38,10 @@ from sglang.multimodal_gen.runtime.server_warmup import (
|
|||||||
run_async_client_warmup,
|
run_async_client_warmup,
|
||||||
should_run_synthetic_server_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.srt.utils.json_response import orjson_response
|
||||||
from sglang.version import __version__
|
from sglang.version import __version__
|
||||||
|
|
||||||
@@ -374,6 +377,7 @@ def create_app(server_args: ServerArgs):
|
|||||||
"""
|
"""
|
||||||
Create and configure the FastAPI application instance.
|
Create and configure the FastAPI application instance.
|
||||||
"""
|
"""
|
||||||
|
globally_suppress_loggers()
|
||||||
app = FastAPI(lifespan=lifespan)
|
app = FastAPI(lifespan=lifespan)
|
||||||
app.add_middleware(
|
app.add_middleware(
|
||||||
CORSMiddleware,
|
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.pipelines_core.schedule_batch import OutputBatch
|
||||||
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
|
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.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||||
from sglang.srt.utils.json_response import orjson_response
|
from sglang.srt.utils.json_response import orjson_response
|
||||||
|
|
||||||
@@ -45,6 +45,26 @@ class DiffusionModelCard(ModelCard):
|
|||||||
pipeline_class: Optional[str] = None
|
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):
|
async def _handle_lora_request(req: Any, success_msg: str, failure_msg: str):
|
||||||
try:
|
try:
|
||||||
output: OutputBatch = await async_scheduler_client.forward(req)
|
output: OutputBatch = await async_scheduler_client.forward(req)
|
||||||
@@ -183,27 +203,7 @@ async def available_models():
|
|||||||
if not server_args:
|
if not server_args:
|
||||||
raise HTTPException(status_code=500, detail="Server args not initialized")
|
raise HTTPException(status_code=500, detail="Server args not initialized")
|
||||||
|
|
||||||
model_info = get_model_info(
|
model_card = _build_model_card(server_args, server_args.model_path)
|
||||||
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)
|
|
||||||
|
|
||||||
# Return dict directly to preserve extended fields (ModelList strips them)
|
# Return dict directly to preserve extended fields (ModelList strips them)
|
||||||
return {"object": "list", "data": [model_card.model_dump()]}
|
return {"object": "list", "data": [model_card.model_dump()]}
|
||||||
@@ -229,24 +229,5 @@ async def retrieve_model(model: str):
|
|||||||
status_code=404,
|
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 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)
|
await MESH_STORE.update_fields(job_id, update_fields)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{e}")
|
logger.exception("Mesh job %s failed", job_id)
|
||||||
await MESH_STORE.update_fields(
|
await MESH_STORE.update_fields(
|
||||||
job_id, {"status": "failed", "error": {"message": str(e)}}
|
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),
|
limit: Optional[int] = Query(None, ge=1, le=100),
|
||||||
order: Optional[str] = Query("desc"),
|
order: Optional[str] = Query("desc"),
|
||||||
):
|
):
|
||||||
order = (order or "desc").lower()
|
jobs = await MESH_STORE.list_page(after=after, limit=limit, order=order)
|
||||||
if order not in ("asc", "desc"):
|
items = [MeshResponse(**job) for job in jobs]
|
||||||
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]
|
|
||||||
return MeshListResponse(data=items)
|
return MeshListResponse(data=items)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,5 +1,5 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any, Dict, Optional
|
||||||
|
|
||||||
|
|
||||||
class AsyncDictStore:
|
class AsyncDictStore:
|
||||||
@@ -36,9 +36,31 @@ class AsyncDictStore:
|
|||||||
async with self._lock:
|
async with self._lock:
|
||||||
return self._items.pop(key, None)
|
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:
|
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
|
# Global stores shared by OpenAI entrypoints
|
||||||
|
|||||||
@@ -489,7 +489,7 @@ async def _dispatch_job_async(
|
|||||||
update_fields.update(final_media_fields)
|
update_fields.update(final_media_fields)
|
||||||
await VIDEO_STORE.update_fields(job_id, update_fields)
|
await VIDEO_STORE.update_fields(job_id, update_fields)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{e}")
|
logger.exception("Video job %s failed", job_id)
|
||||||
await VIDEO_STORE.update_fields(
|
await VIDEO_STORE.update_fields(
|
||||||
job_id,
|
job_id,
|
||||||
{
|
{
|
||||||
@@ -842,25 +842,8 @@ async def list_videos(
|
|||||||
limit: Optional[int] = Query(None, ge=1, le=100),
|
limit: Optional[int] = Query(None, ge=1, le=100),
|
||||||
order: Optional[str] = Query("desc"),
|
order: Optional[str] = Query("desc"),
|
||||||
):
|
):
|
||||||
# Normalize order
|
jobs = await VIDEO_STORE.list_page(after=after, limit=limit, order=order)
|
||||||
order = (order or "desc").lower()
|
items = [VideoResponse(**job) for job in jobs]
|
||||||
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]
|
|
||||||
return VideoListResponse(data=items)
|
return VideoListResponse(data=items)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -66,7 +66,7 @@ async def run_zeromq_broker(server_args: ServerArgs):
|
|||||||
# 1. Receive a request from an offline client
|
# 1. Receive a request from an offline client
|
||||||
payload = await socket.recv()
|
payload = await socket.recv()
|
||||||
request_batch = pickle.loads(payload)
|
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
|
# 2. Forward the request to the main Scheduler via the shared client
|
||||||
response_batch = await async_scheduler_client.forward(request_batch)
|
response_batch = await async_scheduler_client.forward(request_batch)
|
||||||
|
|||||||
Reference in New Issue
Block a user