[diffusion] UX: replace deprecated ORJSONResponse with orjson_response (#21755)
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -9,7 +9,6 @@ from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
from fastapi import APIRouter, FastAPI, Request
|
||||
from fastapi.responses import ORJSONResponse
|
||||
|
||||
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
|
||||
from sglang.multimodal_gen.runtime.entrypoints.openai import image_api, video_api
|
||||
@@ -25,6 +24,7 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import (
|
||||
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
|
||||
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
|
||||
from sglang.version import __version__
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -235,7 +235,7 @@ vertex_router = APIRouter()
|
||||
@vertex_router.post(VERTEX_ROUTE)
|
||||
async def vertex_generate(vertex_req: VertexGenerateReqInput):
|
||||
if not vertex_req.instances:
|
||||
return ORJSONResponse({"predictions": []})
|
||||
return orjson_response({"predictions": []})
|
||||
|
||||
server_args = get_global_server_args()
|
||||
params = vertex_req.parameters or {}
|
||||
@@ -263,7 +263,7 @@ async def vertex_generate(vertex_req: VertexGenerateReqInput):
|
||||
|
||||
results = await asyncio.gather(*futures)
|
||||
|
||||
return ORJSONResponse({"predictions": results})
|
||||
return orjson_response({"predictions": results})
|
||||
|
||||
|
||||
def create_app(server_args: ServerArgs):
|
||||
|
||||
@@ -2,7 +2,6 @@ import time
|
||||
from typing import Any, List, Optional, Union
|
||||
|
||||
from fastapi import APIRouter, Body, HTTPException
|
||||
from fastapi.responses import ORJSONResponse
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from sglang.multimodal_gen.registry import get_model_info
|
||||
@@ -17,6 +16,7 @@ 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.utils.json_response import orjson_response
|
||||
|
||||
router = APIRouter(prefix="/v1")
|
||||
logger = init_logger(__name__)
|
||||
@@ -173,7 +173,7 @@ async def list_loras():
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.get("/models", response_class=ORJSONResponse)
|
||||
@router.get("/models")
|
||||
async def available_models():
|
||||
"""Show available models. OpenAI-compatible endpoint with extended diffusion info."""
|
||||
server_args = get_global_server_args()
|
||||
@@ -206,7 +206,7 @@ async def available_models():
|
||||
return {"object": "list", "data": [model_card.model_dump()]}
|
||||
|
||||
|
||||
@router.get("/models/{model:path}", response_class=ORJSONResponse)
|
||||
@router.get("/models/{model:path}")
|
||||
async def retrieve_model(model: str):
|
||||
"""Retrieve a model instance. OpenAI-compatible endpoint with extended diffusion info."""
|
||||
server_args = get_global_server_args()
|
||||
@@ -214,9 +214,8 @@ async def retrieve_model(model: str):
|
||||
raise HTTPException(status_code=500, detail="Server args not initialized")
|
||||
|
||||
if model != server_args.model_path:
|
||||
return ORJSONResponse(
|
||||
status_code=404,
|
||||
content={
|
||||
return orjson_response(
|
||||
{
|
||||
"error": {
|
||||
"message": f"The model '{model}' does not exist",
|
||||
"type": "invalid_request_error",
|
||||
@@ -224,6 +223,7 @@ async def retrieve_model(model: str):
|
||||
"code": "model_not_found",
|
||||
}
|
||||
},
|
||||
status_code=404,
|
||||
)
|
||||
|
||||
model_info = get_model_info(
|
||||
|
||||
@@ -1,13 +1,13 @@
|
||||
"""Weight update API for the diffusion engine."""
|
||||
|
||||
from fastapi import APIRouter, Request
|
||||
from fastapi.responses import ORJSONResponse
|
||||
|
||||
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
|
||||
GetWeightsChecksumReqInput,
|
||||
UpdateWeightFromDiskReqInput,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client
|
||||
from sglang.srt.utils.json_response import orjson_response
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -18,7 +18,7 @@ async def update_weights_from_disk(request: Request):
|
||||
body = await request.json()
|
||||
model_path = body.get("model_path")
|
||||
if not model_path:
|
||||
return ORJSONResponse(
|
||||
return orjson_response(
|
||||
{"success": False, "message": "model_path is required"},
|
||||
status_code=400,
|
||||
)
|
||||
@@ -32,7 +32,7 @@ async def update_weights_from_disk(request: Request):
|
||||
try:
|
||||
response = await async_scheduler_client.forward(req)
|
||||
except Exception as e:
|
||||
return ORJSONResponse(
|
||||
return orjson_response(
|
||||
{"success": False, "message": str(e)},
|
||||
status_code=500,
|
||||
)
|
||||
@@ -40,7 +40,7 @@ async def update_weights_from_disk(request: Request):
|
||||
result = response.output
|
||||
success = result.get("success", False)
|
||||
message = result.get("message", "Unknown status")
|
||||
return ORJSONResponse(
|
||||
return orjson_response(
|
||||
{"success": success, "message": message},
|
||||
status_code=200 if success else 400,
|
||||
)
|
||||
@@ -57,6 +57,6 @@ async def get_weights_checksum(request: Request):
|
||||
try:
|
||||
response = await async_scheduler_client.forward(req)
|
||||
except Exception as e:
|
||||
return ORJSONResponse({"error": str(e)}, status_code=500)
|
||||
return orjson_response({"error": str(e)}, status_code=500)
|
||||
|
||||
return ORJSONResponse(response.output, status_code=200)
|
||||
return orjson_response(response.output, status_code=200)
|
||||
|
||||
Reference in New Issue
Block a user