Skip to content

Commit 6aada00

Browse files
committed
[feat] Let the OpenAI server plug in a non-CUDA generator
1 parent a28f2ba commit 6aada00

4 files changed

Lines changed: 62 additions & 18 deletions

File tree

fastvideo/entrypoints/openai/api_server.py

Lines changed: 28 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22
# (https://github.com/sgl-project/sglang/blob/main/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py)
33

44
from contextlib import asynccontextmanager
5-
from collections.abc import AsyncIterator
5+
from collections.abc import AsyncIterator, Callable
66
import os
77

88
import uvicorn
@@ -18,7 +18,8 @@
1818
clear_state,
1919
set_state,
2020
)
21-
from fastvideo.entrypoints.openai.serving_engine import OpenAIServingEngine
21+
from fastvideo.entrypoints.openai.serving_engine import OpenAIServingEngine, ServingGenerator
22+
from fastvideo.entrypoints.openai.protocol import VideoGenerationRequest
2223
from fastvideo.entrypoints.video_generator import VideoGenerator
2324
from fastvideo.fastvideo_args import FastVideoArgs
2425
from fastvideo.logger import init_logger
@@ -60,10 +61,13 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
6061
served_model_name: str | None = app.state.served_model_name
6162
default_request: GenerationRequest | None = getattr(app.state, "default_request", None)
6263

63-
logger.info("Loading model from %s ...", args.model_path)
64-
generator = VideoGenerator.from_fastvideo_args(args)
65-
serving_engine = OpenAIServingEngine(generator)
66-
logger.info("Model loaded successfully.")
64+
logger.info("Initializing %s generation runtime for %s ...", app.state.runtime, args.model_path)
65+
# A non-CUDA runtime (e.g. MLX) supplies its own generator_factory instead
66+
# of the default VideoGenerator/Executor path -- see MLXWanGenerator.
67+
factory = app.state.generator_factory
68+
generator = factory() if factory is not None else VideoGenerator.from_fastvideo_args(args)
69+
serving_engine = OpenAIServingEngine(generator, app.state.video_request_validator)
70+
logger.info("Generation runtime ready.")
6771

6872
set_state(
6973
generator,
@@ -91,8 +95,18 @@ def create_app(
9195
output_dir: str = DEFAULT_OUTPUT_DIR,
9296
default_request: GenerationRequest | None = None,
9397
served_model_name: str | None = None,
98+
*,
99+
generator_factory: Callable[[], ServingGenerator] | None = None,
100+
video_request_validator: Callable[[VideoGenerationRequest], None] | None = None,
101+
runtime: str = "cuda",
94102
) -> FastAPI:
95-
"""Build the FastAPI application with all routers mounted"""
103+
"""Build the FastAPI application with all routers mounted.
104+
105+
``generator_factory``/``video_request_validator``/``runtime`` let a
106+
non-CUDA backend (MLX) plug in its own generator and request rules
107+
without a separate app-building path. Defaults reproduce the original
108+
CUDA-only behavior exactly.
109+
"""
96110

97111
app = FastAPI(
98112
title="FastVideo OpenAI-Compatible API",
@@ -103,6 +117,9 @@ def create_app(
103117
app.state.output_dir = output_dir
104118
app.state.default_request = default_request
105119
app.state.served_model_name = served_model_name
120+
app.state.generator_factory = generator_factory
121+
app.state.video_request_validator = video_request_validator
122+
app.state.runtime = runtime
106123

107124
app.add_middleware(
108125
CORSMiddleware,
@@ -148,7 +165,10 @@ async def openai_validation_error(_request: Request, exc: RequestValidationError
148165

149166
app.include_router(common_router)
150167
app.include_router(video_router)
151-
app.include_router(image_router)
168+
# The MLX runtime only wires video-with-audio generation (see
169+
# fastvideo/mlx_runtime/); image generation has no MLX backend yet.
170+
if runtime != "mlx":
171+
app.include_router(image_router)
152172

153173
@app.get("/health")
154174
async def health():

fastvideo/entrypoints/openai/serving_engine.py

Lines changed: 27 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -5,14 +5,28 @@
55

66
import asyncio
77
from collections.abc import Awaitable, Callable
8-
from typing import Any, TypeVar
8+
from typing import Any, Protocol, TypeVar
99

1010
from fastvideo.api.schema import GenerationRequest
11-
from fastvideo.entrypoints.video_generator import VideoGenerator
11+
from fastvideo.entrypoints.openai.protocol import VideoGenerationRequest
1212

1313
_T = TypeVar("_T")
1414

1515

16+
class ServingGenerator(Protocol):
17+
"""The minimal shape OpenAIServingEngine needs from a generator.
18+
19+
VideoGenerator (CUDA/mp/ray) satisfies this structurally already; an MLX
20+
generator (e.g. MLXWanGenerator) implements just these two methods.
21+
"""
22+
23+
def generate(self, request: GenerationRequest) -> Any:
24+
...
25+
26+
def shutdown(self) -> None:
27+
...
28+
29+
1630
class OpenAIServingEngine:
1731
"""Own generator lifecycle and serialize access to its mutable pipeline.
1832
@@ -24,16 +38,24 @@ class OpenAIServingEngine:
2438
the lock without changing the transport contract.
2539
"""
2640

27-
def __init__(self, generator: VideoGenerator) -> None:
41+
def __init__(self,
42+
generator: ServingGenerator,
43+
video_request_validator: Callable[[VideoGenerationRequest], None] | None = None) -> None:
2844
self._generator = generator
45+
self._video_request_validator = video_request_validator
2946
self._generation_lock = asyncio.Lock()
3047
self._closed = False
3148
self._unhealthy_reason: str | None = None
3249

3350
@property
34-
def generator(self) -> VideoGenerator:
51+
def generator(self) -> ServingGenerator:
3552
return self._generator
3653

54+
def validate_video_request(self, request: VideoGenerationRequest) -> None:
55+
"""Runtime-specific request checks (e.g. MLX's supported-field allowlist)."""
56+
if self._video_request_validator is not None:
57+
self._video_request_validator(request)
58+
3759
@property
3860
def closed(self) -> bool:
3961
return self._closed
@@ -127,4 +149,4 @@ async def shutdown(self) -> None:
127149
await asyncio.to_thread(self._generator.shutdown)
128150

129151

130-
__all__ = ["OpenAIServingEngine"]
152+
__all__ = ["OpenAIServingEngine", "ServingGenerator"]

fastvideo/entrypoints/openai/state.py

Lines changed: 4 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -11,21 +11,20 @@
1111

1212
if TYPE_CHECKING:
1313
from fastvideo.api.schema import GenerationRequest
14-
from fastvideo.entrypoints.openai.serving_engine import OpenAIServingEngine
15-
from fastvideo.entrypoints.video_generator import VideoGenerator
14+
from fastvideo.entrypoints.openai.serving_engine import OpenAIServingEngine, ServingGenerator
1615
from fastvideo.fastvideo_args import FastVideoArgs
1716

1817
DEFAULT_OUTPUT_DIR = "outputs"
1918

20-
_generator: VideoGenerator | None = None
19+
_generator: ServingGenerator | None = None
2120
_serving_engine: OpenAIServingEngine | None = None
2221
_fastvideo_args: FastVideoArgs | None = None
2322
_output_dir: str = DEFAULT_OUTPUT_DIR
2423
_served_model_name: str | None = None
2524
_default_request: GenerationRequest | None = None
2625

2726

28-
def get_generator() -> VideoGenerator:
27+
def get_generator() -> ServingGenerator:
2928
"""Return the global VideoGenerator instance (set during startup)."""
3029
assert _generator is not None, "Server not initialized — generator is None"
3130
return _generator
@@ -62,7 +61,7 @@ def get_default_request() -> GenerationRequest | None:
6261

6362

6463
def set_state(
65-
generator: VideoGenerator,
64+
generator: ServingGenerator,
6665
serving_engine: OpenAIServingEngine,
6766
fastvideo_args: FastVideoArgs,
6867
output_dir: str,

fastvideo/entrypoints/openai/video_api.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -331,6 +331,9 @@ async def _parse_video_request(raw_request: Request) -> VideoGenerationRequest:
331331

332332
async def _adapt_request(request_id: str, request: VideoGenerationRequest) -> GenerationRequest:
333333
try:
334+
# Runtime-specific checks (e.g. MLX's supported-field allowlist) run
335+
# before the CUDA-oriented model/LoRA validation below.
336+
get_serving_engine().validate_video_request(request)
334337
validate_model_and_lora(request, get_server_args(), get_served_model_name())
335338
await prepare_reference_media(request_id, request, get_output_dir())
336339
return await asyncio.to_thread(

0 commit comments

Comments
 (0)