Coverage for wrapper/run_httpserver.py: 64%
1210 statements
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
« prev ^ index » next coverage.py v7.15.4, created at 2026-08-09 04:47 +0000
1"""
2REST API server for LMM generation.
4Linux example:
5IMG_BASE64=$(base64 -w 0 benchmark/samples/sample.png)
6cat > payload.json <<EOF
7{
8 "img": "$IMG_BASE64",
9 "prompt": "The person in the right is speaking to the person in the left",
10 "num_frames": 17,
11 "sampling_steps": 5
12}
13EOF
14curl -X POST http://localhost:8080/wan -H "Content-Type: application/json" -d @payload.json
15"""
16import argparse
17import asyncio
18import gzip
19import io
20import json
21import logging
22import mimetypes
23import os
24import pickle
25import sys
26import time
27import errno
28import traceback
29import aiofiles
30import aiofiles.os
32from PIL import Image
34from datetime import datetime
35from datetime import timedelta
36from functools import wraps
37from http import HTTPStatus
38from pickle import UnpicklingError
39from json import JSONDecodeError
41from typing import Tuple
42from typing import Optional
43from typing import Callable
44from typing import Any
45from typing import Awaitable
46from typing import List
47from typing import Dict
48from typing import AsyncIterable
50import torch
51import torch.distributed as dist
52from torch.distributed import DistBackendError
54from hypercorn.config import Config
55from hypercorn.asyncio import serve
57from quart import Quart
58from quart import request
59from quart import jsonify
60from quart import send_file
61from quart import send_from_directory
62from quart import render_template
63from quart import Response
65from wrapper_model import ModelGeneration
67from console_utils import setup_logging
69from media_utils import base64_to_tensor
70from image_utils import img_to_bytesio
71from image_utils import img_to_base64
72from media_utils import get_audio_file_info
73from media_utils import get_image_file_info
74from media_utils import get_text_file_info
75from media_utils import get_video_file_info
76from media_utils import get_tensor_file_info
78import quart_utils
79from quart_utils import QuartReturn
80from quart_utils import get_mime_type
81from quart_utils import get_file_type
82from quart_utils import get_friendly_model_name
83from quart_utils import get_friendly_container_name
84from quart_utils import get_class_emoji
85from quart_utils import format_string
88GPU_SETUP = True
89try:
90 from xfuser import xFuserArgs
91 from xfuser.config import FlexibleArgumentParser
92 from xfuser.config import EngineConfig
93 from xfuser.core.distributed import init_distributed_environment
94except Exception as ex:
95 if "No CUDA GPUs are available" in str(ex):
96 logging.error("No GPUs available, running without xfuser.")
97 elif "Found no NVIDIA driver on your system." in str(ex):
98 logging.error("No NVIDIA driver found, running without xfuser.")
99 elif "cannot import name 'SanaAttnProcessor2_0' from 'diffusers.models.transformers.sana_transformer'" in str(ex):
100 logging.error("Old diffusers version (no SanaAttnProcessor2_0), running without xfuser.")
101 else:
102 logging.error(f"Likely no GPU setup, running without xfuser: {ex}.")
103 logging.error(traceback.format_exc())
104 EngineConfig = None
105 GPU_SETUP = False
107 def init_distributed_environment(rank: int, world_size: int) -> None:
108 """Dummy function when no GPU setup."""
109 logging.warning("No distributed environment setup.")
112# Quart/Flask app configuration
113HOST = "0.0.0.0"
114PORT = 8080
115TMP_DIR = "/tmp"
116app = Quart(__name__)
117route = app.route
118template_filter = app.template_filter
119route_locks: Dict[str, asyncio.Lock] = {}
120last_ping_time = time.time()
122# Models
123models: Dict[str, ModelGeneration] = {}
125EXCLUDED_NCCL_MODELS: List[str] = [
126 "hunyuanimage"
127]
130def get_job_id() -> str:
131 """Generate a unique job ID based on the current timestamp: 20240605T153000123."""
132 return datetime.now().strftime("%Y%m%dT%H%M%S%f")[:-3]
135def get_service_names() -> List[str]:
136 """Get the list of available service names from services.json."""
137 services = {}
138 with open("services.json", "r", encoding="utf-8") as file:
139 services = json.load(file)
140 # Only include entries that represent runnable model services (have a class field)
141 return [name for name, config in services.items() if "class" in config]
144def get_model(
145 service_name: str,
146 sub_module: Optional[str] = None
147) -> Optional[ModelGeneration]:
148 """Get the model instance for the given service name."""
149 if service_name in models:
150 return models[service_name]
151 if sub_module is not None:
152 full_service_name = f"{service_name}{sub_module}"
153 if full_service_name in models:
154 return models[full_service_name]
155 if "mock" in models:
156 return models["mock"]
157 return None
160def exclusive_route(key: str) -> Callable[[Callable[..., Awaitable[Any]]], Callable[..., Awaitable[Any]]]:
161 """Decorator to ensure exclusive access to a route based on a key."""
162 def decorator(func: Callable[..., Awaitable[Any]]) -> Callable[..., Awaitable[Any]]:
163 """Decorator to ensure exclusive access to a route based on a key."""
165 @wraps(func)
166 async def wrapper(*args: Any, **kwargs: Any) -> Any:
167 lock = route_locks.setdefault(key, asyncio.Lock())
168 if lock.locked():
169 request_json = await request.get_json()
170 job_id = request_json.get("job_id")
171 if job_id is not None:
172 logging.warning(f"{job_id} trying to use exclusive route '{key}' already running.")
173 else:
174 logging.warning(f"Exclusive route '{key}' already running.")
175 return jsonify({"error": "Generation in progress"}), HTTPStatus.SERVICE_UNAVAILABLE
176 await lock.acquire()
177 try:
178 return await func(*args, **kwargs)
179 finally:
180 lock.release()
181 return wrapper
183 return decorator
186@template_filter("get_friendly_container_name")
187async def get_friendly_container_name_template(container_name: str) -> str:
188 return await get_friendly_container_name(container_name)
191@template_filter("format_string")
192def format_string_template(value: Any) -> Optional[str]:
193 return format_string(value)
196@template_filter("get_class_emoji")
197async def get_class_emoji_template(container_name: str) -> str:
198 return await get_class_emoji(container_name)
201# HTTP routes
202@route("/", methods=["GET"])
203async def index() -> str:
204 """Render the index HTML page."""
205 models_health = {}
206 for model_name, model in models.items():
207 if model is not None:
208 models_health[model_name] = model.get_health()
209 return await render_template(
210 "index.html",
211 models_health=models_health)
214@route("/health", methods=["GET"])
215async def health() -> Response:
216 """Get health status of all models."""
217 ret = {}
218 for model_name, model in models.items():
219 if model is not None:
220 ret[model_name] = model.get_health()
221 return jsonify(ret)
224@route("/<service_name>/health", methods=["GET"])
225async def model_health(service_name: str) -> QuartReturn:
226 """Get health status of a specific model."""
227 model = get_model(service_name)
228 if model is None:
229 return (
230 {"error": f"{service_name} not initialized"},
231 HTTPStatus.INTERNAL_SERVER_ERROR
232 )
233 return model.get_health()
236@route("/timestamps", methods=["GET"])
237async def timestamps() -> Response:
238 """Get timing timestamps for all models."""
239 ret = {}
240 for model_name, model in models.items():
241 if model is not None:
242 ret[model_name] = model.get_timestamps()
243 return jsonify(ret)
246@route("/files", methods=["GET"])
247async def list_files() -> QuartReturn:
248 """List files in the TMP_DIR directory."""
249 try:
250 files = await quart_utils.list_files(TMP_DIR)
251 return jsonify({"files": files})
252 except Exception as ex:
253 return jsonify({"error": str(ex)}), HTTPStatus.INTERNAL_SERVER_ERROR
256@route("/file/<file_name>", methods=["GET"])
257async def download_file(file_name: str) -> QuartReturn:
258 """Download a file."""
259 filepath = f"{TMP_DIR}/{file_name}"
260 if not await aiofiles.os.path.exists(filepath):
261 return jsonify({"error": "File not found"}), HTTPStatus.NOT_FOUND
263 # Directory listing
264 if await aiofiles.os.path.isdir(filepath):
265 files = await aiofiles.os.listdir(filepath)
266 return jsonify({
267 "files": files
268 })
270 mimetype = get_mime_type(file_name)
271 # Fixes length error when sending large files with send_file, use send_from_directory instead
272 # TODO conditional=True?
273 return await send_from_directory(
274 TMP_DIR,
275 file_name,
276 mimetype=mimetype,
277 as_attachment=True)
280@route("/file_info/<file_name>", methods=["GET"])
281async def file_info(file_name: str) -> QuartReturn:
282 """Get detailed information about a file."""
283 file_path = f"{TMP_DIR}/{file_name}"
284 if not await aiofiles.os.path.exists(file_path):
285 return jsonify({"error": "File not found"}), HTTPStatus.NOT_FOUND
287 file_type = get_file_type(file_name)
288 mimetype, _ = mimetypes.guess_type(file_name)
289 if mimetype is None:
290 mimetype = "application/octet-stream"
292 file_info_ret = {
293 "name": file_name,
294 "size": await aiofiles.os.path.getsize(file_path),
295 "date": await aiofiles.os.path.getmtime(file_path),
296 "type": file_type,
297 "mimetype": mimetype
298 }
300 if file_type == "audio":
301 file_audio_info = get_audio_file_info(file_path)
302 file_info_ret.update(file_audio_info)
303 elif file_type == "video":
304 file_video_info = get_video_file_info(file_path)
305 video_info = file_video_info["video"]
306 file_info_ret.update(video_info)
307 # audio_info = file_video_info["audio"]
308 # file_info_ret.update(audio)
309 elif file_type == "image":
310 file_image_info = get_image_file_info(file_path)
311 file_info_ret.update(file_image_info)
312 elif file_type == "text":
313 file_text_info = get_text_file_info(file_path)
314 file_info_ret.update(file_text_info)
315 elif file_type == "tensor":
316 file_tensor_info = get_tensor_file_info(file_path)
317 file_info_ret.update(file_tensor_info)
319 return jsonify(file_info_ret)
322@route("/yolo", methods=["POST"])
323@exclusive_route("yolo")
324async def yolo_e2e() -> QuartReturn:
325 """YOLO object detection endpoint."""
326 model = get_model("yolo")
327 if not model:
328 return jsonify({"error": "YOLO not initialized"}), HTTPStatus.INTERNAL_SERVER_ERROR
329 try:
330 request_json = await request.get_json()
331 if request_json is None:
332 return jsonify({"error": "No JSON body received"}), HTTPStatus.BAD_REQUEST
334 job_id = request_json.get("job_id") or get_job_id()
335 async with aiofiles.open(f"{TMP_DIR}/{job_id}.json", "w", encoding="utf-8") as file:
336 data = json.dumps(request_json, indent=4)
337 await file.write(data)
338 await file.flush()
340 gen_args = await model.get_rest_args(request_json)
341 args = gen_args["args"]
342 args["job_id"] = job_id
344 if "img" in args:
345 img = args["img"]
346 img.save(f"{TMP_DIR}/{job_id}.png")
348 # Run the one chunk in GPU 0 (blocks for a while)
349 logging.info(f"Generating YOLO with args: {args}")
350 extracted_images = await model.generate(**args)
352 if len(extracted_images) > 0:
353 debug_img = extracted_images[0]
354 debug_img.save(f"{TMP_DIR}/{job_id}_debug.png")
356 images_json = {}
357 for img_id, extracted_image in enumerate(extracted_images[1:]):
358 if not extracted_image:
359 logging.warning(f"No character extracted for image {img_id} in job {job_id}, skipping.")
360 else:
361 extracted_image.save(f"{TMP_DIR}/{job_id}_{img_id:03d}.png")
362 extracted_img_base64 = img_to_base64(extracted_image)
363 images_json[f"image{img_id:03d}"] = extracted_img_base64
364 return jsonify(images_json)
365 except ValueError as value_err:
366 logging.error(f"Error generating YOLO: {value_err}")
367 return jsonify({
368 "error": str(value_err),
369 "traceback": traceback.format_exc()
370 }), HTTPStatus.BAD_REQUEST
371 except Exception as ex:
372 logging.error(f"Error generating YOLO: {ex} {traceback.format_exc()}")
373 return jsonify({
374 "error": str(ex),
375 "traceback": traceback.format_exc()
376 }), HTTPStatus.INTERNAL_SERVER_ERROR
379@route("/podcasttranscript", methods=["POST"])
380async def podcasttranscript_e2e() -> QuartReturn:
381 """podcasttranscript transcript generation endpoint."""
382 model = get_model("podcasttranscript")
383 if not model:
384 return jsonify({"error": "Podcast transcript model not initialized"}), HTTPStatus.INTERNAL_SERVER_ERROR
385 try:
386 request_json = await request.get_json()
387 if request_json is None:
388 return jsonify({"error": "No JSON body received"}), HTTPStatus.BAD_REQUEST
390 job_id = request_json.get("job_id") or get_job_id()
391 async with aiofiles.open(f"{TMP_DIR}/{job_id}.json", "w") as file:
392 data = json.dumps(request_json, indent=4)
393 await file.write(data)
395 gen_args = await model.get_rest_args(request_json)
396 args = gen_args["args"]
397 args["job_id"] = job_id
399 # Run locally in a single GPU
400 logging.info(f"Generating podcast transcript with args: {args}")
401 podcast = await model.generate(**args)
403 if podcast is None:
404 return jsonify({"error": "No podcast generated"}), HTTPStatus.INTERNAL_SERVER_ERROR
406 async with aiofiles.open(f"{TMP_DIR}/{job_id}_podcast.json", "w") as file:
407 data = json.dumps(podcast.model_dump(), indent=4)
408 await file.write(data)
410 return jsonify(podcast.model_dump())
411 except ValueError as value_err:
412 logging.error(f"Error generating podcast transcript: {value_err}")
413 return jsonify({"error": str(value_err)}), HTTPStatus.BAD_REQUEST
414 except Exception as ex:
415 logging.error(f"Error generating podcast transcript: {ex}")
416 return jsonify({
417 "error": str(ex),
418 "traceback": traceback.format_exc()
419 }), HTTPStatus.INTERNAL_SERVER_ERROR
422@route("/slidetranscript/stream", methods=["POST"])
423async def slidetranscript_stream() -> QuartReturn:
424 """Stream slide transcript generation endpoint."""
425 model = get_model("slidetranscript")
426 if not model:
427 return jsonify({"error": "Slide transcript model not initialized"}), HTTPStatus.INTERNAL_SERVER_ERROR
428 try:
429 request_json = await request.get_json()
430 if request_json is None:
431 return jsonify({"error": "No JSON body received"}), HTTPStatus.BAD_REQUEST
433 job_id = request_json.get("job_id") or get_job_id()
434 request_json["job_id"] = job_id
435 async with aiofiles.open(f"{TMP_DIR}/{job_id}.json", "w") as file:
436 data = json.dumps(request_json, indent=4)
437 await file.write(data)
439 gen_args = await model.get_rest_args(request_json)
440 args = gen_args["args"]
441 args["job_id"] = job_id
443 async def slide_transcript_generate() -> AsyncIterable[str]:
444 try:
445 async with aiofiles.open(f"{TMP_DIR}/{job_id}_slides.jsonl", "w") as f_out:
446 async for scene in model.generate_stream(**args):
447 line = json.dumps(scene) + "\n"
448 await f_out.write(line)
449 await f_out.flush()
450 yield line
451 except Exception as ex:
452 logging.exception(f"Exception streaming (job_id={job_id}): {ex}")
453 yield json.dumps({
454 "error": "Stream generation failed",
455 "details": str(ex)
456 }) + "\n"
458 return Response(
459 slide_transcript_generate(),
460 mimetype="application/x-ndjson")
461 except Exception as ex:
462 logging.error(f"Error generating slide transcript: {ex}")
463 return jsonify({
464 "error": str(ex),
465 "traceback": traceback.format_exc()
466 }), HTTPStatus.INTERNAL_SERVER_ERROR
469@route("/podcasttranscript/stream", methods=["POST"])
470async def podcasttranscript_stream() -> QuartReturn:
471 """Stream podcast transcript generation endpoint."""
472 model = get_model("podcasttranscript")
473 if not model:
474 return jsonify({"error": "Podcast transcript model not initialized"}), HTTPStatus.INTERNAL_SERVER_ERROR
475 try:
476 request_json = await request.get_json()
477 if request_json is None:
478 return jsonify({"error": "No JSON body received"}), HTTPStatus.BAD_REQUEST
480 job_id = request_json.get("job_id") or get_job_id()
481 request_json["job_id"] = job_id
482 async with aiofiles.open(f"{TMP_DIR}/{job_id}.json", "w") as file:
483 data = json.dumps(request_json, indent=4)
484 await file.write(data)
486 gen_args = await model.get_rest_args(request_json)
487 args = gen_args["args"]
488 args["job_id"] = job_id
490 async def podcast_transcript_generate() -> AsyncIterable[str]:
491 try:
492 async with aiofiles.open(f"{TMP_DIR}/{job_id}_podcast.jsonl", "w") as f_out:
493 async for scene in model.generate_stream(**args):
494 line = json.dumps(scene) + "\n"
495 await f_out.write(line)
496 await f_out.flush()
497 yield line
498 except Exception as ex:
499 logging.exception(f"Exception streaming (job_id={job_id}): {ex}")
500 yield json.dumps({
501 "error": "Stream generation failed",
502 "details": str(ex)
503 }) + "\n"
505 return Response(
506 podcast_transcript_generate(),
507 mimetype="application/x-ndjson")
508 except ValueError as value_err:
509 logging.error(f"Error generating podcast transcript: {value_err}")
510 return jsonify({"error": str(value_err)}), HTTPStatus.BAD_REQUEST
511 except Exception as ex:
512 logging.error(f"Error generating podcast transcript: {ex}")
513 return jsonify({
514 "error": str(ex),
515 "traceback": traceback.format_exc()
516 }), HTTPStatus.INTERNAL_SERVER_ERROR
519async def gen_video(model: Optional[ModelGeneration]) -> QuartReturn:
520 """Generic video generation endpoint."""
521 if not model:
522 return jsonify({"error": "Not initialized"}), HTTPStatus.INTERNAL_SERVER_ERROR
523 if model.status != "ok":
524 return jsonify({"error": f"Model not ready: {model.status}"}), HTTPStatus.SERVICE_UNAVAILABLE
525 # TODO not all models are constrained to running a request at a time
526 if model.running:
527 logging.warning("Generation in progress.")
528 return jsonify({"error": "Generation in progress"}), HTTPStatus.SERVICE_UNAVAILABLE
529 try:
530 request_json = await request.get_json()
531 if request_json is None:
532 return jsonify({"error": "No JSON body received"}), HTTPStatus.BAD_REQUEST
534 job_id = request_json.get("job_id") or get_job_id()
535 request_json["job_id"] = job_id
536 async with aiofiles.open(f"{TMP_DIR}/{job_id}.json", "w") as file:
537 data = json.dumps(request_json, indent=4)
538 await file.write(data)
540 gen_args = await model.get_rest_args(request_json)
541 args = gen_args["args"]
542 args["job_id"] = job_id
543 args["output_type"] = "video_path"
545 # Trigger video generation in all GPUs
546 await send_task(gen_args)
548 # Run the one chunk in GPU 0 (blocks for a while)
549 logging.info(f"Generating video with args: {args}")
550 video_path = await model.generate(**args)
552 if not video_path:
553 return jsonify({"error": "No video generated"}), HTTPStatus.INTERNAL_SERVER_ERROR
554 if not await aiofiles.os.path.exists(video_path):
555 return jsonify({"error": f"Video file not found: {video_path}"}), HTTPStatus.INTERNAL_SERVER_ERROR
557 return await send_file(
558 video_path,
559 mimetype="video/mp4",
560 as_attachment=True,
561 attachment_filename=f"{job_id}.mp4")
562 except JSONDecodeError as json_err:
563 logging.error(f"Error processing JSON request: {json_err}")
564 return jsonify({"error": "Invalid JSON format"}), HTTPStatus.BAD_REQUEST
565 except ValueError as value_err:
566 logging.error(f"Error generating video: {value_err}")
567 return jsonify({"error": str(value_err)}), HTTPStatus.BAD_REQUEST
568 except Exception as ex:
569 logging.error(f"Error generating video: {ex}")
570 return jsonify({
571 "error": str(ex),
572 "traceback": traceback.format_exc()
573 }), HTTPStatus.INTERNAL_SERVER_ERROR
576async def gen_audio(model: Optional[ModelGeneration]) -> QuartReturn:
577 """Generic audio generation endpoint."""
578 if not model:
579 return jsonify({"error": "Not initialized"}), HTTPStatus.INTERNAL_SERVER_ERROR
580 if model.status != "ok":
581 return jsonify({"error": f"Model not ready: {model.status}"}), HTTPStatus.SERVICE_UNAVAILABLE
582 # TODO not all models are constrained to running a request at a time
583 if model.running:
584 logging.warning("Generation in progress.")
585 return jsonify({"error": "Generation in progress"}), HTTPStatus.SERVICE_UNAVAILABLE
586 try:
587 request_json = await request.get_json()
588 job_id = request_json.get("job_id") or get_job_id()
589 request_json["job_id"] = job_id
590 async with aiofiles.open(f"{TMP_DIR}/{job_id}.json", "w") as file:
591 data = json.dumps(request_json, indent=4)
592 await file.write(data)
594 gen_args = await model.get_rest_args(request_json)
595 args = gen_args["args"]
596 args["job_id"] = job_id
598 # Trigger image generation in all GPUs
599 await send_task(gen_args)
601 # Run the one chunk in GPU 0 (blocks for a while)
602 logging.info(f"Generating audio with args: {args}")
603 audio_path = await model.generate(**args)
605 if not audio_path:
606 return jsonify({"error": "No audio generated"}), HTTPStatus.INTERNAL_SERVER_ERROR
607 if not await aiofiles.os.path.exists(audio_path):
608 return jsonify({"error": f"Audio file not found: {audio_path}"}), HTTPStatus.INTERNAL_SERVER_ERROR
610 return await send_from_directory(
611 os.path.dirname(audio_path),
612 os.path.basename(audio_path),
613 mimetype="audio/wav",
614 as_attachment=True,
615 attachment_filename=f"{job_id}.wav")
616 except JSONDecodeError as json_err:
617 logging.error(f"Error processing JSON request: {json_err}")
618 return jsonify({"error": "Invalid JSON format"}), HTTPStatus.BAD_REQUEST
619 except ValueError as value_err:
620 logging.error(f"Error generating audio: {value_err}")
621 return jsonify({"error": str(value_err)}), HTTPStatus.BAD_REQUEST
622 except Exception as ex:
623 logging.error(f"Error generating audio: {ex}")
624 return jsonify({
625 "error": str(ex),
626 "traceback": traceback.format_exc()
627 }), HTTPStatus.INTERNAL_SERVER_ERROR
630async def gen_img(model: Optional[ModelGeneration]) -> QuartReturn:
631 """Generic image generation endpoint."""
632 if not model:
633 return jsonify({"error": "Not initialized"}), HTTPStatus.INTERNAL_SERVER_ERROR
634 if model.status != "ok":
635 return jsonify({"error": f"Model not ready: {model.status}"}), HTTPStatus.SERVICE_UNAVAILABLE
636 # TODO not all models are constrained to running a request at a time
637 if model.running:
638 logging.warning("Generation in progress.")
639 return jsonify({"error": "Generation in progress"}), HTTPStatus.SERVICE_UNAVAILABLE
640 try:
641 request_json = await request.get_json()
642 job_id = request_json.get("job_id") or get_job_id()
643 async with aiofiles.open(f"{TMP_DIR}/{job_id}.json", mode="w") as file:
644 data = json.dumps(request_json, indent=4)
645 await file.write(data)
647 gen_args = await model.get_rest_args(request_json)
648 args = gen_args["args"]
649 args["job_id"] = job_id
651 # Trigger image generation in all GPUs
652 await send_task(gen_args)
654 # Run the one chunk in GPU 0 (blocks for a while)
655 img_wrapper = await model.generate(**args)
657 if isinstance(img_wrapper, list) and len(img_wrapper) > 0:
658 img = img_wrapper[0]
659 else:
660 img = img_wrapper
662 if not img:
663 return jsonify({"error": "No image generated"}), HTTPStatus.INTERNAL_SERVER_ERROR
664 if not isinstance(img, Image.Image):
665 return jsonify({"error": f"No image generated: {type(img)}"}), HTTPStatus.INTERNAL_SERVER_ERROR
667 img.save(f"{TMP_DIR}/{job_id}.png")
668 img_io = img_to_bytesio(img)
669 if img_io is None:
670 return jsonify({"error": "Failed to convert image to bytes"}), HTTPStatus.INTERNAL_SERVER_ERROR
671 return await send_file(
672 img_io,
673 mimetype="image/png",
674 as_attachment=True,
675 attachment_filename=f"{job_id}.png")
676 except JSONDecodeError as json_err:
677 logging.error(f"Error processing JSON request: {json_err}")
678 return jsonify({"error": "Invalid JSON format"}), HTTPStatus.BAD_REQUEST
679 except ValueError as value_err:
680 logging.error(f"Error generating image: {value_err}")
681 return jsonify({"error": str(value_err)}), HTTPStatus.BAD_REQUEST
682 except Exception as ex:
683 logging.error(f"Error generating image: {ex}")
684 return jsonify({
685 "error": str(ex),
686 "traceback": traceback.format_exc()
687 }), HTTPStatus.INTERNAL_SERVER_ERROR
690@route("/fantasytalking", methods=["POST"])
691@exclusive_route("fantasytalking")
692async def fantasytalking_e2e() -> QuartReturn:
693 """FantasyTalking video generation endpoint."""
694 model = get_model("fantasytalking")
695 return await gen_video(model)
698@route("/kokoro", methods=["POST"])
699@exclusive_route("kokoro")
700async def kokoro_e2e() -> QuartReturn:
701 """Kokoro audio generation endpoint."""
702 model = get_model("kokoro")
703 return await gen_audio(model)
706@route("/dia", methods=["POST"])
707@exclusive_route("dia")
708async def dia_e2e() -> QuartReturn:
709 """DIA audio generation endpoint."""
710 model = get_model("dia")
711 return await gen_audio(model)
714@route("/xtts", methods=["POST"])
715@exclusive_route("xtts")
716async def xtts_e2e() -> QuartReturn:
717 """XTTS audio generation endpoint."""
718 model = get_model("xtts")
719 return await gen_audio(model)
722@route("/thinksound", methods=["POST"])
723@exclusive_route("thinksound")
724async def thinksound_e2e() -> QuartReturn:
725 """ThinkSound audio generation endpoint."""
726 model = get_model("thinksound")
727 return await gen_audio(model)
730@route("/vibevoice", methods=["POST"])
731@exclusive_route("vibevoice")
732async def vibevoice_e2e() -> QuartReturn:
733 """VibeVoice audio generation endpoint."""
734 model = get_model("vibevoice")
735 return await gen_audio(model)
738@route("/flux", methods=["POST"])
739@exclusive_route("flux")
740async def flux_e2e() -> QuartReturn:
741 """Flux image generation endpoint."""
742 model = get_model("flux")
743 return await gen_img(model)
746@route("/fluxupscaler", methods=["POST"])
747@exclusive_route("fluxupscaler")
748async def fluxupscaler_e2e() -> QuartReturn:
749 """Flux image upscaling endpoint."""
750 model = get_model("fluxupscaler")
751 return await gen_img(model)
754@route("/fluxupscaler/video", methods=["POST"])
755@exclusive_route("fluxupscaler")
756async def fluxupscaler_video_e2e() -> QuartReturn:
757 """Flux video upscaling endpoint."""
758 model = get_model("fluxupscaler")
759 return await gen_video(model)
762@route("/fluxkontext", methods=["POST"])
763@exclusive_route("fluxkontext")
764async def fluxkontext_e2e() -> QuartReturn:
765 """Flux Kontext image generation endpoint."""
766 model = get_model("fluxkontext")
767 return await gen_img(model)
770@route("/fluxkrea", methods=["POST"])
771@exclusive_route("fluxkrea")
772async def fluxkrea_e2e() -> QuartReturn:
773 """Flux Krea image generation endpoint."""
774 model = get_model("fluxkrea")
775 return await gen_img(model)
778@route("/flux2", methods=["POST"])
779@exclusive_route("flux2")
780async def flux2_e2e() -> QuartReturn:
781 """FLUX.2-dev image generation endpoint."""
782 model = get_model("flux2")
783 return await gen_img(model)
786@route("/flux2klein", methods=["POST"])
787@exclusive_route("flux2klein")
788async def flux2klein_e2e() -> QuartReturn:
789 """FLUX.2-klein-9B image generation endpoint."""
790 model = get_model("flux2klein")
791 return await gen_img(model)
794@route("/hidream", methods=["POST"])
795@exclusive_route("hidream")
796async def hidream_e2e() -> QuartReturn:
797 """HiDream image generation endpoint."""
798 model = get_model("hidream")
799 return await gen_img(model)
802@route("/qwenimage", methods=["POST"])
803@exclusive_route("qwenimage")
804async def qwenimage_e2e() -> QuartReturn:
805 """QwenImage image generation endpoint."""
806 model = get_model("qwenimage")
807 return await gen_img(model)
810@route("/qwenimageedit", methods=["POST"])
811@exclusive_route("qwenimageedit")
812async def qwenimageedit_e2e() -> QuartReturn:
813 """QwenImageEdit image generation endpoint."""
814 model = get_model("qwenimageedit")
815 return await gen_img(model)
818@route("/hunyuanimage", methods=["POST"])
819@exclusive_route("hunyuanimage")
820async def hunyuanimage_e2e() -> QuartReturn:
821 """Hunyuan Image generation endpoint."""
822 model = get_model("hunyuanimage")
823 return await gen_img(model)
826@route("/januspro", methods=["POST"])
827@exclusive_route("januspro")
828async def januspro_e2e() -> QuartReturn:
829 """JanusPro image generation endpoint."""
830 model = get_model("januspro")
831 return await gen_img(model)
834@route("/llamagen", methods=["POST"])
835@exclusive_route("llamagen")
836async def llamagen_e2e() -> QuartReturn:
837 """LlamaGen image generation endpoint."""
838 model = get_model("llamagen")
839 return await gen_img(model)
842@route("/cogview", methods=["POST"])
843@exclusive_route("cogview")
844async def cogview_e2e() -> QuartReturn:
845 """CogView4 image generation endpoint."""
846 model = get_model("cogview")
847 return await gen_img(model)
850@route("/bagel", methods=["POST"])
851@exclusive_route("bagel")
852async def bagel_e2e() -> QuartReturn:
853 """Bagel image generation endpoint."""
854 model = get_model("bagel")
855 return await gen_img(model)
858@route("/imageresize", methods=["POST"])
859async def imageresize_e2e() -> QuartReturn:
860 """Basic pillow image resize endpoint."""
861 model = get_model("imageresize")
862 return await gen_img(model)
865@route("/4kagent", methods=["POST"])
866@exclusive_route("4kagent")
867async def fourk_agent_e2e() -> QuartReturn:
868 """4KAgent image super-resolution endpoint."""
869 model = get_model("4kagent")
870 return await gen_img(model)
873@route("/realesrgan", methods=["POST"])
874@exclusive_route("realesrgan")
875async def realesrgan_e2e() -> QuartReturn:
876 """Real-ESRGAN image upscaling endpoint."""
877 model = get_model("realesrgan")
878 return await gen_img(model)
881@route("/realesrgan/video", methods=["POST"])
882@exclusive_route("realesrgan")
883async def realesrgan_video_e2e() -> QuartReturn:
884 model = get_model("realesrgan")
885 return await gen_video(model)
888"""
889# TODO add a multi form data endpoint for realesrgan
890@route("/realesrgan/video", methods=["POST"])
891@exclusive_route("realesrgan")
892async def realesrgan_video_e2e():
893 form = await request.form
894 file = (await request.files)["video"]
896 job_id = form.get("job_id")
897 width = int(form.get("width", 768))
898 height = int(form.get("height", 576))
900 print(f"Received job_id={job_id}, width={width}, height={height}, filename={file.filename}")
902 # Simulate reading video file
903 video_bytes = await file.read()
904 print(f"Received video size: {len(video_bytes)} bytes")
906 # Simulate processing and return dummy video
907 dummy_output = io.BytesIO(video_bytes) # Just echo input for demo
908 dummy_output.seek(0)
909 return await send_file(
910 dummy_output,
911 mimetype="video/mp4",
912 download_name="upscaled.mp4",
913 )
914"""
917@route("/hunyuanframepack", methods=["POST"])
918@exclusive_route("hunyuanframepack")
919async def hunyuanframepack_e2e() -> QuartReturn:
920 """Hunyuan FramePack video generation endpoint."""
921 model = get_model("hunyuanframepack")
922 return await gen_video(model)
925@route("/hunyuanframepackf1", methods=["POST"])
926@exclusive_route("hunyuanframepack")
927async def hunyuanframepackf1_e2e() -> QuartReturn:
928 """Hunyuan FramePack F1 video generation endpoint."""
929 model = get_model("hunyuanframepackf1")
930 return await gen_video(model)
933@route("/hunyuanframepack/vae", methods=["POST"])
934@exclusive_route("hunyuanframepack")
935async def hunyuanframepack_vae() -> QuartReturn:
936 """Hunyuan FramePack VAE decode endpoint."""
937 model = get_model("hunyuanframepack", "vae")
938 if model is None or model.vae is None:
939 return jsonify({"error": "Hunyuan FramePack VAE not available"}), HTTPStatus.INTERNAL_SERVER_ERROR
940 return await gen_video(model)
943@route("/hunyuanframepackvae", methods=["POST"])
944@exclusive_route("hunyuanframepack")
945async def hunyuanframepackvae() -> QuartReturn:
946 """Hunyuan FramePack VAE decode endpoint."""
947 return await hunyuanframepack_vae()
950@route("/hunyuanframepack/vae/<job_id>", methods=["POST"])
951@exclusive_route("hunyuanframepack")
952async def hunyuanframepack_vae_binary(job_id: str) -> QuartReturn:
953 """Hunyuan FramePack VAE decode binary endpoint."""
954 model = get_model("hunyuanframepack", "vae")
955 if model is None or model.vae is None:
956 return jsonify({"error": "Hunyuan FramePack VAE not available"}), HTTPStatus.INTERNAL_SERVER_ERROR
957 if job_id is None:
958 return jsonify({"error": "Missing 'job_id' parameter"}), HTTPStatus.BAD_REQUEST
959 try:
960 encoding = request.headers.get('Content-Encoding', '').lower()
961 data = await request.get_data()
962 if not isinstance(data, (bytes, bytearray)):
963 return jsonify({"error": f"Invalid data: {type(data)}"}), HTTPStatus.BAD_REQUEST
964 data_bytes = io.BytesIO(data)
965 if encoding == 'gzip':
966 with gzip.GzipFile(fileobj=data_bytes, mode='rb') as file_gzip:
967 data_decompressed = file_gzip.read()
968 data_bytes = io.BytesIO(data_decompressed)
969 decompressed_len = len(data_decompressed)
970 else:
971 decompressed_len = len(data)
972 latents = torch.load(data_bytes, weights_only=True)
974 logging.info(
975 f"Process latent for '{job_id}' with {decompressed_len} bytes (HTTP:{len(data)}) "
976 f"shape {latents.shape} and {latents.dtype}.")
977 torch.save(latents, f"{TMP_DIR}/{job_id}_latents.pt")
978 request_json = {
979 "job_id": job_id,
980 "size": decompressed_len,
981 "http_size": len(data),
982 "shape": str(latents.shape),
983 "dtype": str(latents.dtype),
984 "filename": f"{TMP_DIR}/{job_id}_latents.pt",
985 }
986 async with aiofiles.open(f"{TMP_DIR}/{job_id}.json", mode="w") as file:
987 data = json.dumps(request_json, indent=4)
988 await file.write(data)
990 args = {
991 "latents": latents,
992 "job_id": job_id,
993 "output_type": "video_path"
994 }
996 # Run the one chunk in GPU 0 (blocks for a while)
997 logging.info(f"Decoding latents with args: {args}")
998 video_path = await model.generate(**args)
1000 return await send_file(
1001 video_path,
1002 mimetype="video/mp4",
1003 as_attachment=True,
1004 attachment_filename=f"{job_id}.mp4")
1005 except ValueError as value_err:
1006 logging.error(f"Error processing VAE latents: {value_err}")
1007 return jsonify({
1008 "error": str(value_err),
1009 "traceback": traceback.format_exc()
1010 }), HTTPStatus.BAD_REQUEST
1011 except Exception as ex:
1012 logging.error(f"Error processing VAE latents: {ex}")
1013 return jsonify({
1014 "error": str(ex),
1015 "traceback": traceback.format_exc()
1016 }), HTTPStatus.INTERNAL_SERVER_ERROR
1019@route("/wan", methods=["POST"])
1020@exclusive_route("wan")
1021async def wan_e2e() -> QuartReturn:
1022 """Wan video generation endpoint."""
1023 model = get_model("wan")
1024 return await gen_video(model)
1027@route("/wan/vae", methods=["POST"])
1028@exclusive_route("wan")
1029async def wan_vae() -> QuartReturn:
1030 """Wan VAE decode endpoint."""
1031 # TODO this is not tested
1032 model = get_model("wanvae")
1033 if model is None or model.vae is None:
1034 return jsonify({"error": "Wan VAE not available"}), HTTPStatus.INTERNAL_SERVER_ERROR
1036 try:
1037 request_json = await request.get_json()
1038 job_id = request_json.get("job_id") or get_job_id()
1039 async with aiofiles.open(f"{TMP_DIR}/{job_id}.json", mode="w") as file:
1040 data = json.dumps(request_json, indent=4)
1041 await file.write(data)
1043 latents_base64 = request_json.get("latents", None)
1044 if latents_base64 is None:
1045 return jsonify({"error": "Missing 'latents' parameter"}), HTTPStatus.BAD_REQUEST
1046 latents = base64_to_tensor(latents_base64)
1048 # Run locally in a single GPU
1049 pixels = await asyncio.to_thread(model.vae_decode, latents)
1051 # Save the pixels to a file
1052 file_path = f"{TMP_DIR}/{job_id}_pixels.pt"
1053 torch.save(pixels, file_path)
1055 return await send_file(
1056 file_path,
1057 mimetype="application/octet-stream",
1058 as_attachment=True,
1059 attachment_filename=f"{job_id}_pixels.pt"
1060 )
1061 except JSONDecodeError as json_err:
1062 logging.error(f"Error processing JSON request: {json_err}")
1063 return jsonify({"error": "Invalid JSON format"}), HTTPStatus.BAD_REQUEST
1064 except ValueError as value_err:
1065 logging.error(f"Error processing VAE latents: {value_err}")
1066 return jsonify({
1067 "error": str(value_err),
1068 "traceback": traceback.format_exc()
1069 }), HTTPStatus.BAD_REQUEST
1070 except Exception as ex:
1071 logging.error(f"Error processing VAE latents: {ex}")
1072 return jsonify({
1073 "error": str(ex),
1074 "traceback": traceback.format_exc()
1075 }), HTTPStatus.INTERNAL_SERVER_ERROR
1078@route("/wanvae", methods=["POST"])
1079@exclusive_route("wan")
1080async def wanvae() -> QuartReturn:
1081 """Wan VAE decode endpoint."""
1082 return await wan_vae()
1085@route("/wan22", methods=["POST"])
1086@exclusive_route("wan22")
1087async def wan22_e2e() -> QuartReturn:
1088 """Wan 2.2 video generation endpoint."""
1089 model = get_model("wan22")
1090 return await gen_video(model)
1093@route("/hunyuanavatar", methods=["POST"])
1094@exclusive_route("hunyuanavatar")
1095async def hunyuanavatar_e2e() -> QuartReturn:
1096 """HunyuanAvatar video generation endpoint."""
1097 model = get_model("hunyuanavatar")
1098 return await gen_video(model)
1101@route("/ltx", methods=["POST"])
1102@exclusive_route("ltx")
1103async def ltx_e2e() -> QuartReturn:
1104 """LTX video generation endpoint."""
1105 model = get_model("ltx")
1106 return await gen_video(model)
1109@route("/longcatvideo", methods=["POST"])
1110@exclusive_route("longcatvideo")
1111async def longcatvideo_e2e() -> QuartReturn:
1112 """LongCat-Video generation endpoint."""
1113 model = get_model("longcatvideo")
1114 return await gen_video(model)
1117@route("/interrupt", methods=["POST"])
1118async def interrupt_gen() -> QuartReturn:
1119 """Interrupt any ongoing generation."""
1120 for model in models.values():
1121 if model:
1122 model.interrupt()
1123 return jsonify({"status": "interrupted"}), HTTPStatus.OK
1126# Distributed environment setup
1127rank = 0
1128local_rank = 0
1129device = torch.device(f"cuda:{rank}")
1130node_rank = 0
1131world_size = 1
1132local_world_size = 1
1134MAX_PAYLOAD_BYTES = 1 * 1024 * 1024 * 1024 # 1 GB max payload size
1136# Signal file path for distributed tasks, this only works in single node multi-GPU mode
1137SIGNAL_0_PATH = f"{TMP_DIR}/streamwise_signal.txt"
1140def setup_dist_environment() -> None:
1141 """
1142 Initialize the distributed environment if running in multi-GPU mode.
1143 If `param_only` is True, only set the parameters without initializing the process group.
1144 """
1145 if "MASTER_ADDR" not in os.environ:
1146 logging.info("MASTER_ADDR not set, setting local.")
1147 os.environ["MASTER_ADDR"] = "localhost"
1148 os.environ["MASTER_PORT"] = "12355"
1149 os.environ["RANK"] = "0"
1150 os.environ["WORLD_SIZE"] = "1"
1152 if not torch.cuda.is_available():
1153 logging.warning("CUDA not available, skipping distributed initialization.")
1154 return
1156 global rank
1157 global local_rank
1158 global node_rank
1159 global world_size
1160 global local_world_size
1161 rank = int(os.getenv("RANK", 0))
1162 local_rank = int(os.getenv("LOCAL_RANK", 0))
1163 node_rank = int(os.getenv("NODE_RANK", 0))
1164 world_size = int(os.getenv("WORLD_SIZE", 1))
1165 local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE") or os.environ.get(
1166 "NPROC_PER_NODE") or torch.cuda.device_count())
1168 # With MIG, the device plugin restricts CUDA_VISIBLE_DEVICES to only the allocated
1169 # MIG instance(s) for this process, so the valid device indices start at 0.
1170 # Use the number of visible devices to clamp local_rank to a valid index.
1171 num_visible_devices = torch.cuda.device_count()
1172 device_id = local_rank if local_rank < num_visible_devices else 0
1173 torch.cuda.set_device(device_id)
1175 if world_size > num_visible_devices:
1176 logging.warning(
1177 f"world_size={world_size} but only {num_visible_devices} visible CUDA device(s). "
1178 "This usually means the container is running on a MIG partition and "
1179 "torchrun was invoked with too many processes. "
1180 "Clamping world_size to the number of visible devices."
1181 )
1182 world_size = num_visible_devices
1183 local_world_size = num_visible_devices
1185 logging.info(f"[{rank}] Initializing distributed: "
1186 f"rank={rank}, local_rank={local_rank}, node_rank={node_rank}, "
1187 f"world_size={world_size}, local_world_size={local_world_size}, "
1188 f"device={device_id}")
1191def init_dist_environment() -> None:
1192 """Initialize the distributed process group."""
1193 if not torch.cuda.is_available():
1194 logging.debug(f"[{rank}] CUDA not available, skipping distributed initialization.")
1195 return
1196 if dist.is_initialized():
1197 logging.info(f"[{rank}] Distributed process group already initialized.")
1198 return
1200 dist.init_process_group(
1201 backend="nccl",
1202 init_method="env://",
1203 rank=rank,
1204 world_size=world_size,
1205 timeout=timedelta(hours=24), # Prevent NCCL timeout
1206 )
1208 if dist.get_world_size() > 1:
1209 init_distributed_environment(
1210 rank=dist.get_rank(),
1211 world_size=dist.get_world_size()
1212 )
1214 if not dist.is_initialized():
1215 raise RuntimeError("Distributed process group not initialized.")
1218async def wait_for_everybody() -> None:
1219 """
1220 Wait for all processes to reach this point.
1221 This is useful to ensure that all workers are ready before sending tasks.
1222 This only works in single node multi-GPU mode.
1223 """
1224 if not dist.is_initialized() or world_size <= 1:
1225 return
1227 logging.info(f"[{rank}] Waiting for all {local_world_size} workers to be ready...")
1229 # Specify that we are ready
1230 signal_worker_path = f"{TMP_DIR}/streamwise_signal_worker_{local_rank:03d}.txt"
1231 async with aiofiles.open(signal_worker_path, mode="w") as file:
1232 await file.write(f"ready at {time.strftime('%Y-%m-%d %H:%M:%S', time.localtime())}")
1234 # Wait for the signal file from all other ranks
1235 for other_rank in range(local_world_size):
1236 other_signal_worker_path = f"{TMP_DIR}/streamwise_signal_worker_{other_rank:03d}.txt"
1237 while not await aiofiles.os.path.exists(other_signal_worker_path):
1238 await asyncio.sleep(0.1) # Wait until the file exists
1240 async with aiofiles.open(other_signal_worker_path, mode="r") as file:
1241 content = await file.read()
1242 logging.info(f"[{rank}] Worker {other_rank}: {content}.")
1244 logging.info(f"[{rank}] All {local_world_size} workers are ready.")
1247async def send_task(gen_task: dict) -> None:
1248 """
1249 Send a generation task to all workers in the distributed environment (through NCCL).
1250 """
1251 if not dist.is_initialized():
1252 logging.warning(f"[{rank}] Torch distributed not initialized.")
1253 return
1254 if rank != 0:
1255 logging.error(f"[{rank}] Task can only be sent from rank 0.")
1256 return
1257 if world_size <= 1:
1258 logging.debug(f"[{rank}] Single GPU mode, skipping task broadcast.")
1259 return
1260 task_id = gen_task.get("task", "")
1261 if task_id in EXCLUDED_NCCL_MODELS:
1262 logging.debug(f"[{rank}] No NCCL-based parallelism for {task_id}.")
1263 return
1265 global last_ping_time
1266 last_ping_time = time.time()
1268 try:
1269 payload_bytes = await asyncio.to_thread(pickle.dumps, gen_task)
1270 payload_buffer = bytearray(payload_bytes)
1271 payload_tensor = torch.frombuffer(payload_buffer, dtype=torch.uint8).to("cuda")
1272 payload_size = torch.tensor([payload_tensor.numel()], dtype=torch.int64, device="cuda")
1274 if payload_size.item() > MAX_PAYLOAD_BYTES:
1275 logging.error(f"[{rank}] Payload too large: {payload_size.item()} bytes.")
1276 return
1278 # Notify that rank 0 is ready
1279 async with aiofiles.open(SIGNAL_0_PATH, mode="w") as file:
1280 await file.write(f"ready at {time.strftime('%Y-%m-%d %H:%M:%S', time.localtime())}")
1282 # Broadcast size and data
1283 logging.info(f"[{rank}] Broadcasting payload with {payload_size.item()} bytes.")
1284 dist.broadcast(payload_size, src=0)
1285 torch.cuda.synchronize()
1286 dist.barrier()
1288 # Clean up the rank 0 signal file
1289 await asyncio.to_thread(os.remove, SIGNAL_0_PATH)
1291 dist.broadcast(payload_tensor, src=0)
1292 torch.cuda.synchronize()
1294 logging.debug(f"[{rank}] Broadcast complete.")
1295 except DistBackendError as dist_err:
1296 logging.error(f"[{rank}] Cannot send task NCCL error: {dist_err}")
1297 except Exception as ex:
1298 logging.exception(f"[{rank}] Cannot send task {type(ex)}: {ex}")
1301async def nccl_worker() -> None:
1302 """Torch requests (through NCCL) from rank 0."""
1303 if rank == 0:
1304 logging.error(f"[{rank}] Worker should not be started on rank 0.")
1305 return
1307 while dist.is_initialized():
1308 logging.info(f"[{rank}] Waiting for tasks from rank 0...")
1309 try:
1310 # Wait until signal file exists
1311 while not await aiofiles.os.path.exists(SIGNAL_0_PATH):
1312 await asyncio.sleep(0.1)
1314 # Receive payload size
1315 payload_size = torch.tensor([0], dtype=torch.int64, device="cuda")
1316 dist.broadcast(payload_size, src=0)
1317 torch.cuda.synchronize()
1318 dist.barrier()
1320 payload_size_val = int(payload_size.item())
1321 if payload_size_val <= 0 or payload_size_val > MAX_PAYLOAD_BYTES:
1322 logging.error(f"[{rank}] Invalid payload size: {payload_size_val}.")
1323 continue
1325 # Receive payload data
1326 payload_tensor = torch.empty(payload_size_val, dtype=torch.uint8, device="cuda")
1327 dist.broadcast(payload_tensor, src=0)
1328 torch.cuda.synchronize()
1330 payload_bytes = payload_tensor.cpu().numpy().tobytes()
1331 payload = pickle.loads(payload_bytes) # nosec B301 - internal IPC from rank 0
1333 if not isinstance(payload, dict):
1334 logging.error(f"[{rank}] Invalid payload received: {payload}.")
1335 continue
1337 if not payload:
1338 logging.error(f"[{rank}] Empty payload received.")
1339 continue
1341 if "task" not in payload:
1342 logging.error(f"[{rank}] No 'task' in payload: {payload}.")
1343 continue
1345 gen_task = payload
1347 task_id = gen_task["task"]
1348 if task_id == "ping":
1349 logging.debug(f"[{rank}] Received ping task to keep NCCL alive.")
1350 elif models.get(task_id) is None:
1351 logging.error(f"[{rank}] Model '{task_id}' not initialized.")
1352 else:
1353 model = models[task_id]
1354 args = gen_task.get("args", {})
1355 if rank > 1:
1356 logging.info(f"[{rank}] Work to do for {task_id}: {len(args)} arguments.")
1357 else:
1358 logging.info(f"[{rank}] Work to do for {task_id}:")
1359 for key, value in args.items():
1360 if isinstance(value, list) and len(value) > 5:
1361 logging.info(f"[{rank}] {key}: {value[0:5]}... ({len(value)} items)")
1362 else:
1363 logging.info(f"[{rank}] {key}: {value}")
1364 # This can block a little
1365 await model.generate(**args)
1366 except UnpicklingError as pickle_err:
1367 logging.error(f"[{rank}] Parsing message: {pickle_err}.")
1368 except ValueError as value_err:
1369 logging.error(f"[{rank}] Processing task: {value_err}.")
1370 except DistBackendError as dist_err:
1371 logging.error(f"[{rank}] Processing task NCCL error: {dist_err}.")
1372 except Exception as ex:
1373 logging.error(f"[{rank}] Processing task: {ex}.", exc_info=True)
1375 await asyncio.sleep(0.1) # Breathing time between requests just in case
1376 logging.info(f"[{rank}] Exiting worker thread.")
1379def is_model_running() -> bool:
1380 """Check if any model is currently running a generation task."""
1381 for model in models.values():
1382 if model.running:
1383 return True
1384 return False
1387def get_model_names() -> List[str]:
1388 """Get the list of available model names."""
1389 return list(models.keys())
1392def is_nccl_excluded_model() -> bool:
1393 """Check if any of the loaded models are in the NCCL excluded list."""
1394 model_names = get_model_names()
1395 for model_name in EXCLUDED_NCCL_MODELS:
1396 if model_name in model_names:
1397 return True
1398 return False
1401def arg_parsing() -> Tuple[argparse.Namespace, Optional[EngineConfig]]:
1402 """Parse command line arguments and return the parsed args and engine config."""
1403 if GPU_SETUP:
1404 parser = FlexibleArgumentParser(description="REST API for LMM generation")
1405 else:
1406 parser = argparse.ArgumentParser(description="REST API for LMM generation")
1408 parser.add_argument("--host", type=str, default=HOST, help="Host to bind the server")
1409 parser.add_argument("--port", type=int, default=PORT, help="Port to bind the server")
1410 parser.add_argument("--certfile", type=str, default=None, help="Path to SSL certificate file for HTTPS")
1411 parser.add_argument("--keyfile", type=str, default=None, help="Path to SSL private key file for HTTPS")
1413 parser.add_argument("--wan", action="store_true", help="Wan 2.1 model")
1414 parser.add_argument("--wan21", action="store_true", help="Wan 2.1 model")
1415 parser.add_argument("--wanvae", action="store_true", help="Wan 2.1 VAE model")
1416 parser.add_argument("--wan_variation", choices=["480p", "720p"],
1417 default="480p", help="Wan 2.1 model resolution (480p or 720p)")
1418 parser.add_argument("--wan22", action="store_true", help="Wan 2.2 model")
1419 parser.add_argument("--hunyuanframepack", action="store_true", help="Hunyuan FramePack model")
1420 parser.add_argument("--hunyuanframepackf1", action="store_true", help="Hunyuan FramePack F1 model")
1421 parser.add_argument("--hunyuanframepackvae", action="store_true", help="Hunyuan FramePack VAE model")
1422 parser.add_argument("--flux", action="store_true", help="Flux model")
1423 parser.add_argument("--fluxupscaler", action="store_true", help="Flux Upscaler model")
1424 parser.add_argument("--fluxkontext", action="store_true", help="Flux Kontext model")
1425 parser.add_argument("--fluxkrea", action="store_true", help="Flux Krea model")
1426 parser.add_argument("--flux2", action="store_true", help="FLUX.2-dev model")
1427 parser.add_argument("--flux2klein", action="store_true", help="FLUX.2-klein-9B model")
1428 parser.add_argument("--cogview", action="store_true", help="CogView4 model")
1429 parser.add_argument("--hidream", action="store_true", help="HiDream model")
1430 parser.add_argument("--qwenimage", action="store_true", help="Qwen Image model")
1431 parser.add_argument("--qwenimageedit", action="store_true", help="Qwen Image Edit model")
1432 parser.add_argument("--hunyuanimage", action="store_true", help="Hunyuan Image model")
1433 parser.add_argument("--januspro", action="store_true", help="Janus Pro model")
1434 parser.add_argument("--llamagen", action="store_true", help="LlamaGen model")
1435 parser.add_argument("--kokoro", action="store_true", help="Kokoro model")
1436 parser.add_argument("--dia", action="store_true", help="Dia model")
1437 parser.add_argument("--xtts", action="store_true", help="XTTS model")
1438 parser.add_argument("--vibevoice", action="store_true", help="VibeVoice model")
1439 parser.add_argument("--thinksound", action="store_true", help="ThinkSound model")
1440 parser.add_argument("--fantasytalking", action="store_true", help="Fantasy Talking model")
1441 parser.add_argument("--podcasttranscript", action="store_true", help="Podcast transcript wrapper")
1442 parser.add_argument("--slidetranscript", action="store_true", help="Slide transcript wrapper")
1443 parser.add_argument("--yolo", action="store_true", help="YOLO model")
1444 parser.add_argument("--bagel", action="store_true", help="Bagel model")
1445 parser.add_argument("--imageresize", action="store_true", help="Image Resize")
1446 parser.add_argument("--realesrgan", action="store_true", help="Real-ESRGAN model")
1447 parser.add_argument("--4kagent", action="store_true", help="4K Agent")
1448 parser.add_argument("--hunyuanavatar", action="store_true", help="Hunyuan-Avatar model")
1449 parser.add_argument("--ltx", action="store_true", help="LTX-Video model")
1450 parser.add_argument("--longcatvideo", action="store_true", help="LongCat-Video model")
1451 parser.add_argument("--mock", action="store_true", help="Mock model")
1453 if GPU_SETUP and world_size > 1:
1454 args = xFuserArgs.add_cli_args(parser).parse_args()
1455 engine_args = xFuserArgs.from_cli_args(args)
1456 engine_config, _ = engine_args.create_config()
1457 else:
1458 # Add others for compatibility with xFuserArgs
1459 parser.add_argument("--ulysses_degree", type=int, default=1, help="Ulysses degree")
1460 parser.add_argument("--ring_degree", type=int, default=16, help="Ring degree")
1461 parser.add_argument("--use_torch_compile", action="store_true", help="Use torch.compile if available")
1462 args = parser.parse_args()
1463 engine_config = None
1465 return args, engine_config
1468async def load_model_wrapper_file(rank: int, model_name: str) -> None:
1469 """
1470 Set the model wrapper file path based on the model name.
1471 For example: /wan/wrapper_wan21.py
1472 """
1473 friendly_name = await get_friendly_model_name(model_name)
1474 if rank == 0:
1475 logging.info(f"[{rank}] Loading {friendly_name} model.")
1477 if await aiofiles.os.path.exists(f"/wrapper/{model_name}/wrapper_{model_name}.py"):
1478 sys.path.append(f"/wrapper/{model_name}")
1479 return
1480 if await aiofiles.os.path.exists(f"/{model_name}/wrapper_{model_name}.py"):
1481 sys.path.append(f"/{model_name}")
1482 return
1483 if await aiofiles.os.path.exists(f"{model_name}/wrapper_{model_name}.py"):
1484 sys.path.append(f"{model_name}")
1485 return
1486 if await aiofiles.os.path.exists(f"wrapper/{model_name}/wrapper_{model_name}.py"):
1487 sys.path.append(f"wrapper/{model_name}")
1488 return
1489 if await aiofiles.os.path.exists(f"wrapper_{model_name}.py"):
1490 sys.path.append(".")
1491 return
1493 if rank == 0:
1494 files = await aiofiles.os.listdir("/")
1495 logging.info(f"[{rank}] Model wrapper not found. Files available: {files}.")
1496 raise FileNotFoundError(f"Model {friendly_name} not found")
1499async def init_model(
1500 args: argparse.Namespace,
1501 engine_config: Optional[EngineConfig],
1502) -> None:
1503 """Initialize models based on parsed arguments."""
1504 if args.wan or args.wan21:
1505 model_name = "wan"
1506 await load_model_wrapper_file(rank, model_name)
1507 from wrapper_wan21 import Wan21VideoGeneration
1508 wan_ckpt_dir = "/wan/Wan2.1-I2V-14B-480P"
1509 if args.wan_variation == "720p":
1510 wan_ckpt_dir = "/wan/Wan2.1-I2V-14B-720P"
1511 models[model_name] = Wan21VideoGeneration(
1512 ckpt_dir=wan_ckpt_dir,
1513 engine_config=engine_config,
1514 )
1516 if args.wanvae:
1517 model_name = "wanvae"
1518 await load_model_wrapper_file(rank, model_name)
1519 # TODO implement WanVideoVAEGeneration
1521 if args.wan22:
1522 model_name = "wan22"
1523 await load_model_wrapper_file(rank, model_name)
1524 from wrapper_wan22 import Wan22VideoGeneration
1525 wan_ckpt_dir = "/wan/Wan2.2-I2V-A14B"
1526 models[model_name] = Wan22VideoGeneration(
1527 model_name=model_name,
1528 ckpt_dir=wan_ckpt_dir,
1529 engine_config=engine_config)
1531 if args.hunyuanframepack:
1532 model_name = "hunyuanframepack"
1533 await load_model_wrapper_file(rank, model_name)
1534 from wrapper_hunyuanframepack import HunyuanFramepackGeneration
1535 models[model_name] = HunyuanFramepackGeneration(engine_config=engine_config)
1537 if args.hunyuanframepackf1:
1538 model_name = "hunyuanframepackf1"
1539 await load_model_wrapper_file(rank, model_name)
1540 from wrapper_hunyuanframepackf1 import HunyuanFramepackF1Generation
1541 models[model_name] = HunyuanFramepackF1Generation(engine_config=engine_config)
1543 if args.hunyuanframepackvae:
1544 model_name = "hunyuanframepackvae"
1545 await load_model_wrapper_file(rank, model_name)
1546 from wrapper_hunyuanframepackvae import HunyuanFramepackVAEGeneration
1547 models[model_name] = HunyuanFramepackVAEGeneration()
1549 if args.flux:
1550 model_name = "flux"
1551 await load_model_wrapper_file(rank, model_name)
1552 from wrapper_flux import FluxGeneration
1553 models[model_name] = FluxGeneration(engine_config=engine_config)
1555 if args.fluxupscaler:
1556 model_name = "fluxupscaler"
1557 await load_model_wrapper_file(rank, model_name)
1558 from wrapper_fluxupscaler import FluxUpscalerGeneration
1559 models[model_name] = FluxUpscalerGeneration(engine_config=engine_config)
1561 if args.fluxkontext:
1562 model_name = "fluxkontext"
1563 await load_model_wrapper_file(rank, model_name)
1564 from wrapper_fluxkontext import FluxKontextGeneration
1565 models[model_name] = FluxKontextGeneration(engine_config=engine_config)
1567 if args.fluxkrea:
1568 model_name = "fluxkrea"
1569 await load_model_wrapper_file(rank, model_name)
1570 from wrapper_fluxkrea import FluxKreaGeneration
1571 models[model_name] = FluxKreaGeneration(engine_config=engine_config)
1573 if args.flux2:
1574 model_name = "flux2"
1575 await load_model_wrapper_file(rank, model_name)
1576 from wrapper_flux2 import Flux2Generation
1577 models[model_name] = Flux2Generation(engine_config=engine_config)
1579 if args.flux2klein:
1580 model_name = "flux2klein"
1581 await load_model_wrapper_file(rank, model_name)
1582 from wrapper_flux2klein import Flux2KleinGeneration
1583 models[model_name] = Flux2KleinGeneration(engine_config=engine_config)
1585 if args.cogview:
1586 model_name = "cogview"
1587 await load_model_wrapper_file(rank, model_name)
1588 from wrapper_cogview import CogViewGeneration
1589 models[model_name] = CogViewGeneration(engine_config=engine_config)
1591 if args.hidream:
1592 model_name = "hidream"
1593 await load_model_wrapper_file(rank, model_name)
1594 from wrapper_hidream import HiDreamGeneration
1595 models[model_name] = HiDreamGeneration(engine_config=engine_config)
1597 if args.qwenimage:
1598 model_name = "qwenimage"
1599 await load_model_wrapper_file(rank, model_name)
1600 from wrapper_qwenimage import QwenImageGeneration
1601 models[model_name] = QwenImageGeneration(engine_config=engine_config)
1603 if args.qwenimageedit:
1604 model_name = "qwenimageedit"
1605 await load_model_wrapper_file(rank, model_name)
1606 from wrapper_qwenimageedit import QwenImageEditGeneration
1607 models[model_name] = QwenImageEditGeneration(engine_config=engine_config)
1609 if args.hunyuanimage:
1610 model_name = "hunyuanimage"
1611 await load_model_wrapper_file(rank, model_name)
1612 from wrapper_hunyuanimage import HunyuanImageGeneration
1613 models[model_name] = HunyuanImageGeneration(engine_config=engine_config)
1615 if args.januspro:
1616 model_name = "januspro"
1617 await load_model_wrapper_file(rank, model_name)
1618 from wrapper_januspro import JanusProGeneration
1619 models[model_name] = JanusProGeneration(engine_config=engine_config)
1621 if args.llamagen:
1622 model_name = "llamagen"
1623 await load_model_wrapper_file(rank, model_name)
1624 from wrapper_llamagen import LlamaGenGeneration
1625 models[model_name] = LlamaGenGeneration(engine_config=engine_config)
1627 if args.kokoro:
1628 model_name = "kokoro"
1629 await load_model_wrapper_file(rank, model_name)
1630 from wrapper_kokoro import KokoroGeneration
1631 models[model_name] = KokoroGeneration()
1633 if args.dia:
1634 model_name = "dia"
1635 await load_model_wrapper_file(rank, model_name)
1636 from wrapper_dia import DiaGeneration
1637 models[model_name] = DiaGeneration()
1639 if args.xtts:
1640 model_name = "xtts"
1641 await load_model_wrapper_file(rank, model_name)
1642 from wrapper_xtts import XTTSGeneration
1643 models[model_name] = XTTSGeneration()
1645 if args.vibevoice:
1646 model_name = "vibevoice"
1647 await load_model_wrapper_file(rank, model_name)
1648 from wrapper_vibevoice import VibeVoiceGeneration
1649 models[model_name] = VibeVoiceGeneration()
1651 if args.thinksound:
1652 model_name = "thinksound"
1653 await load_model_wrapper_file(rank, model_name)
1654 from wrapper_thinksound import ThinkSoundGeneration
1655 models[model_name] = ThinkSoundGeneration()
1657 if args.fantasytalking:
1658 model_name = "fantasytalking"
1659 await load_model_wrapper_file(rank, model_name)
1660 from wrapper_fantasytalking import FantasyTalking
1661 models[model_name] = FantasyTalking(engine_config=engine_config)
1663 if args.hunyuanavatar:
1664 model_name = "hunyuanavatar"
1665 await load_model_wrapper_file(rank, model_name)
1666 from wrapper_hunyuanavatar import HunyuanAvatarGeneration
1667 models[model_name] = HunyuanAvatarGeneration(engine_config=engine_config)
1669 if args.podcasttranscript:
1670 model_name = "podcasttranscript"
1671 await load_model_wrapper_file(rank, model_name)
1672 from wrapper_podcasttranscript import PodcastTranscriptGenerator
1673 models[model_name] = PodcastTranscriptGenerator()
1675 if args.slidetranscript:
1676 model_name = "slidetranscript"
1677 await load_model_wrapper_file(rank, model_name)
1678 from wrapper_slidetranscript import SlideTranscriptGenerator
1679 models[model_name] = SlideTranscriptGenerator()
1681 if args.yolo:
1682 model_name = "yolo"
1683 await load_model_wrapper_file(rank, model_name)
1684 from wrapper_yolo import ImageCharacterExtractor
1685 models[model_name] = ImageCharacterExtractor()
1687 if args.bagel:
1688 model_name = "bagel"
1689 await load_model_wrapper_file(rank, model_name)
1690 from wrapper_bagel import BagelGeneration
1691 models[model_name] = BagelGeneration()
1693 if args.imageresize:
1694 model_name = "imageresize"
1695 await load_model_wrapper_file(rank, model_name)
1696 from wrapper_imageresize import ImageResize
1697 models[model_name] = ImageResize()
1699 if args.realesrgan:
1700 model_name = "realesrgan"
1701 await load_model_wrapper_file(rank, model_name)
1702 from wrapper_realesrgan import RealESRGANGeneration
1703 models[model_name] = RealESRGANGeneration()
1705 if getattr(args, "4kagent"): # Starting with number has issues
1706 model_name = "4kagent"
1707 await load_model_wrapper_file(rank, model_name)
1708 from wrapper_4kagent import Upscale4KAgent
1709 models[model_name] = Upscale4KAgent()
1711 if args.ltx:
1712 model_name = "ltx"
1713 await load_model_wrapper_file(rank, model_name)
1714 from wrapper_ltx import LTXVideoGeneration
1715 models[model_name] = LTXVideoGeneration()
1717 if args.longcatvideo:
1718 model_name = "longcatvideo"
1719 await load_model_wrapper_file(rank, model_name)
1720 from wrapper_longcatvideo import LongCatVideoGeneration
1721 models[model_name] = LongCatVideoGeneration()
1723 if args.mock:
1724 model_name = "mock"
1725 await load_model_wrapper_file(rank, model_name)
1726 from wrapper_mock import MockGeneration
1727 models[model_name] = MockGeneration()
1729 for model_name, model in models.items():
1730 try:
1731 logging.info(f"[{rank}] Initializing model '{model_name}'...")
1732 await asyncio.to_thread(model.init)
1733 except Exception as ex:
1734 logging.error(f"[{rank}] Error during '{model_name}' initialization: {ex}.")
1735 try:
1736 # Run async warmup in separate process with proper event loop handling
1737 logging.info(f"[{rank}] Warming up model '{model_name}'...")
1738 await model.warmup()
1739 except torch.OutOfMemoryError as oom_err:
1740 logging.error(f"[{rank}] OOM during '{model_name}' warmup: {oom_err}.")
1741 except Exception as ex:
1742 err_msg = str(ex)
1743 logging.error(f"[{rank}] Error during '{model_name}' warmup: {err_msg}.")
1744 if "Model not initialized." not in err_msg:
1745 logging.error(traceback.format_exc())
1746 if rank == 0:
1747 logging.info(f"[{rank}] Model '{model_name}' loaded.")
1750async def run_httpserver(
1751 host: str = HOST,
1752 port: int = PORT,
1753 certfile: Optional[str] = None,
1754 keyfile: Optional[str] = None,
1755) -> None:
1756 """HTTP/HTTPS server runs in the main process (rank 0)."""
1757 config = Config()
1758 config.bind = [f"{host}:{port}"]
1760 config.accesslog = "-"
1762 # Increase max request body size to 128 MB (default is 16 MB)
1763 config.limit_max_request_size = 128 * 1024 * 1024 # type: ignore[attr-defined]
1764 config.wsgi_max_body_size = 128 * 1024 * 1024
1765 app.config["MAX_CONTENT_LENGTH"] = 128 * 1024 * 1024
1767 # Configure for better concurrency (allow more concurrent connections)
1768 config.worker_connections = 64 # type: ignore[attr-defined]
1769 config.keep_alive_timeout = 5 * 60 # Keep connections alive for 5 minutes
1770 config.graceful_timeout = 30 # Graceful shutdown timeout
1772 if certfile:
1773 config.certfile = certfile
1774 if keyfile:
1775 config.keyfile = keyfile
1777 scheme = "https" if certfile else "http"
1778 logging.info(f"[{rank}] Starting {scheme.upper()} server on {scheme}://{host}:{port} with routes:")
1779 for rule in app.url_map.iter_rules():
1780 logging.info(f"[{rank}] - {rule.rule}")
1782 logging.debug(f"[{rank}] HTTP server config:")
1783 for key, value in app.config.items():
1784 logging.debug(f"{key}: {value}")
1786 try:
1787 await serve(app, config)
1788 except OSError as os_err:
1789 if os_err.errno == errno.EADDRINUSE:
1790 logging.error(f"{host}:{port} already in use.")
1791 raise
1794async def main() -> None:
1795 """
1796 Main entry point for the application.
1797 It starts:
1798 - The HTTP server in the main process (rank 0).
1799 - The NCCL workers in other processes (rank > 0).
1800 """
1801 try:
1802 setup_dist_environment()
1804 args, engine_config = arg_parsing()
1806 if not args.hunyuanavatar:
1807 init_dist_environment() # Hunyuan-Avatar handles its own distributed environment
1809 if rank == 0:
1810 # HTTP server runs in the main process (rank 0)
1811 http_task = asyncio.create_task(run_httpserver(
1812 host=args.host,
1813 port=args.port,
1814 certfile=args.certfile,
1815 keyfile=args.keyfile,
1816 ))
1818 await init_model(args, engine_config)
1820 await wait_for_everybody()
1822 # Keep alive ping for NCCL workers
1823 try:
1824 SLEEP_TIME_PING_SECONDS = 60.0
1825 while http_task is not None and not http_task.done():
1826 await asyncio.sleep(SLEEP_TIME_PING_SECONDS)
1827 if world_size > 1 and not is_model_running() and not is_nccl_excluded_model():
1828 logging.info(f"[{rank}] Sending ping to workers.")
1829 args_ping = {
1830 "task": "ping",
1831 "args": {}
1832 }
1833 await send_task(args_ping)
1834 finally:
1835 if http_task is not None and not http_task.done():
1836 http_task.cancel()
1837 await http_task
1838 else:
1839 await init_model(args, engine_config)
1841 await wait_for_everybody()
1843 # Start the workers in the other GPUs
1844 await nccl_worker()
1845 finally:
1846 if dist.is_initialized():
1847 dist.destroy_process_group()
1848 logging.info(f"[{rank}] Exiting main process.")
1851if __name__ == "__main__":
1852 setup_logging(
1853 path=TMP_DIR,
1854 file_name="streamwise.log",
1855 level=logging.INFO)
1857 try:
1858 asyncio.run(main())
1859 except OSError as os_err:
1860 if os_err.errno == errno.EADDRINUSE:
1861 logging.error(f"Port already in use: {os_err}")
1862 else:
1863 raise
1864 except Exception as ex:
1865 logging.error(f"Fatal error in main: {ex}", exc_info=True)