Add v1/models endpoint to diffusion model APIs so that they can be discovered by model gateway (#16425)

This commit is contained in:
Kangyan-Zhou
2026-01-06 14:28:24 -08:00
committed by GitHub
parent 3271e0e76d
commit 18e2ef09d7
4 changed files with 185 additions and 4 deletions

View File

@@ -55,9 +55,15 @@ async def health():
return {"status": "ok"}
@health_router.get("/models")
@health_router.get("/models", deprecated=True)
async def get_models(request: Request):
"""Get information about the model served by this server."""
"""
Get information about the model served by this server.
.. deprecated::
Use /v1/models instead for OpenAI-compatible model discovery.
This endpoint will be removed in a future version.
"""
from sglang.multimodal_gen.registry import get_model_info
server_args: ServerArgs = request.app.state.server_args

View File

@@ -1,7 +1,9 @@
from typing import Any, Optional
from fastapi import APIRouter, Body, HTTPException
from fastapi.responses import ORJSONResponse
from sglang.multimodal_gen.registry import get_model_info
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
MergeLoraWeightsReq,
SetLoraReq,
@@ -11,11 +13,23 @@ from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBa
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.utils.logging_utils import init_logger
from sglang.srt.entrypoints.openai.protocol import ModelCard
router = APIRouter(prefix="/v1")
logger = init_logger(__name__)
class DiffusionModelCard(ModelCard):
"""Extended ModelCard with diffusion-specific fields."""
num_gpus: Optional[int] = None
task_type: Optional[str] = None
dit_precision: Optional[str] = None
vae_precision: Optional[str] = None
pipeline_name: Optional[str] = None
pipeline_class: Optional[str] = None
async def _handle_lora_request(req: Any, success_msg: str, failure_msg: str):
try:
output: OutputBatch = await async_scheduler_client.forward(req)
@@ -117,3 +131,71 @@ async def model_info():
"model_path": server_args.model_path,
}
return result
@router.get("/models", response_class=ORJSONResponse)
async def available_models():
"""Show available models. OpenAI-compatible endpoint with extended diffusion info."""
server_args = get_global_server_args()
if not server_args:
raise HTTPException(status_code=500, detail="Server args not initialized")
model_info = get_model_info(server_args.model_path)
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 {"object": "list", "data": [model_card.model_dump()]}
@router.get("/models/{model:path}", response_class=ORJSONResponse)
async def retrieve_model(model: str):
"""Retrieve a model instance. OpenAI-compatible endpoint with extended diffusion info."""
server_args = get_global_server_args()
if not server_args:
raise HTTPException(status_code=500, detail="Server args not initialized")
if model != server_args.model_path:
return ORJSONResponse(
status_code=404,
content={
"error": {
"message": f"The model '{model}' does not exist",
"type": "invalid_request_error",
"param": "model",
"code": "model_not_found",
}
},
)
model_info = get_model_info(server_args.model_path)
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()