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