Add v1/models endpoint to diffusion model APIs so that they can be discovered by model gateway (#16425)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user