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

1""" 

2REST API server for LMM generation. 

3 

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 

31 

32from PIL import Image 

33 

34from datetime import datetime 

35from datetime import timedelta 

36from functools import wraps 

37from http import HTTPStatus 

38from pickle import UnpicklingError 

39from json import JSONDecodeError 

40 

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 

49 

50import torch 

51import torch.distributed as dist 

52from torch.distributed import DistBackendError 

53 

54from hypercorn.config import Config 

55from hypercorn.asyncio import serve 

56 

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 

64 

65from wrapper_model import ModelGeneration 

66 

67from console_utils import setup_logging 

68 

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 

77 

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 

86 

87 

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 

106 

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.") 

110 

111 

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() 

121 

122# Models 

123models: Dict[str, ModelGeneration] = {} 

124 

125EXCLUDED_NCCL_MODELS: List[str] = [ 

126 "hunyuanimage" 

127] 

128 

129 

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] 

133 

134 

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] 

142 

143 

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 

158 

159 

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.""" 

164 

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 

182 

183 return decorator 

184 

185 

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) 

189 

190 

191@template_filter("format_string") 

192def format_string_template(value: Any) -> Optional[str]: 

193 return format_string(value) 

194 

195 

196@template_filter("get_class_emoji") 

197async def get_class_emoji_template(container_name: str) -> str: 

198 return await get_class_emoji(container_name) 

199 

200 

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) 

212 

213 

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) 

222 

223 

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() 

234 

235 

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) 

244 

245 

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 

254 

255 

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 

262 

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 }) 

269 

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) 

278 

279 

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 

286 

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" 

291 

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 } 

299 

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) 

318 

319 return jsonify(file_info_ret) 

320 

321 

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 

333 

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() 

339 

340 gen_args = await model.get_rest_args(request_json) 

341 args = gen_args["args"] 

342 args["job_id"] = job_id 

343 

344 if "img" in args: 

345 img = args["img"] 

346 img.save(f"{TMP_DIR}/{job_id}.png") 

347 

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) 

351 

352 if len(extracted_images) > 0: 

353 debug_img = extracted_images[0] 

354 debug_img.save(f"{TMP_DIR}/{job_id}_debug.png") 

355 

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 

377 

378 

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 

389 

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) 

394 

395 gen_args = await model.get_rest_args(request_json) 

396 args = gen_args["args"] 

397 args["job_id"] = job_id 

398 

399 # Run locally in a single GPU 

400 logging.info(f"Generating podcast transcript with args: {args}") 

401 podcast = await model.generate(**args) 

402 

403 if podcast is None: 

404 return jsonify({"error": "No podcast generated"}), HTTPStatus.INTERNAL_SERVER_ERROR 

405 

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) 

409 

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 

420 

421 

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 

432 

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) 

438 

439 gen_args = await model.get_rest_args(request_json) 

440 args = gen_args["args"] 

441 args["job_id"] = job_id 

442 

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" 

457 

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 

467 

468 

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 

479 

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) 

485 

486 gen_args = await model.get_rest_args(request_json) 

487 args = gen_args["args"] 

488 args["job_id"] = job_id 

489 

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" 

504 

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 

517 

518 

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 

533 

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) 

539 

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" 

544 

545 # Trigger video generation in all GPUs 

546 await send_task(gen_args) 

547 

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) 

551 

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 

556 

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 

574 

575 

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) 

593 

594 gen_args = await model.get_rest_args(request_json) 

595 args = gen_args["args"] 

596 args["job_id"] = job_id 

597 

598 # Trigger image generation in all GPUs 

599 await send_task(gen_args) 

600 

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) 

604 

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 

609 

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 

628 

629 

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) 

646 

647 gen_args = await model.get_rest_args(request_json) 

648 args = gen_args["args"] 

649 args["job_id"] = job_id 

650 

651 # Trigger image generation in all GPUs 

652 await send_task(gen_args) 

653 

654 # Run the one chunk in GPU 0 (blocks for a while) 

655 img_wrapper = await model.generate(**args) 

656 

657 if isinstance(img_wrapper, list) and len(img_wrapper) > 0: 

658 img = img_wrapper[0] 

659 else: 

660 img = img_wrapper 

661 

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 

666 

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 

688 

689 

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) 

696 

697 

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) 

704 

705 

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) 

712 

713 

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) 

720 

721 

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) 

728 

729 

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) 

736 

737 

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) 

744 

745 

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) 

752 

753 

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) 

760 

761 

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) 

768 

769 

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) 

776 

777 

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) 

784 

785 

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) 

792 

793 

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) 

800 

801 

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) 

808 

809 

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) 

816 

817 

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) 

824 

825 

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) 

832 

833 

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) 

840 

841 

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) 

848 

849 

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) 

856 

857 

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) 

863 

864 

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) 

871 

872 

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) 

879 

880 

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) 

886 

887 

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"] 

895 

896 job_id = form.get("job_id") 

897 width = int(form.get("width", 768)) 

898 height = int(form.get("height", 576)) 

899 

900 print(f"Received job_id={job_id}, width={width}, height={height}, filename={file.filename}") 

901 

902 # Simulate reading video file 

903 video_bytes = await file.read() 

904 print(f"Received video size: {len(video_bytes)} bytes") 

905 

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""" 

915 

916 

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) 

923 

924 

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) 

931 

932 

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) 

941 

942 

943@route("/hunyuanframepackvae", methods=["POST"]) 

944@exclusive_route("hunyuanframepack") 

945async def hunyuanframepackvae() -> QuartReturn: 

946 """Hunyuan FramePack VAE decode endpoint.""" 

947 return await hunyuanframepack_vae() 

948 

949 

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) 

973 

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) 

989 

990 args = { 

991 "latents": latents, 

992 "job_id": job_id, 

993 "output_type": "video_path" 

994 } 

995 

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) 

999 

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 

1017 

1018 

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) 

1025 

1026 

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 

1035 

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) 

1042 

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) 

1047 

1048 # Run locally in a single GPU 

1049 pixels = await asyncio.to_thread(model.vae_decode, latents) 

1050 

1051 # Save the pixels to a file 

1052 file_path = f"{TMP_DIR}/{job_id}_pixels.pt" 

1053 torch.save(pixels, file_path) 

1054 

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 

1076 

1077 

1078@route("/wanvae", methods=["POST"]) 

1079@exclusive_route("wan") 

1080async def wanvae() -> QuartReturn: 

1081 """Wan VAE decode endpoint.""" 

1082 return await wan_vae() 

1083 

1084 

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) 

1091 

1092 

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) 

1099 

1100 

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) 

1107 

1108 

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) 

1115 

1116 

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 

1124 

1125 

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 

1133 

1134MAX_PAYLOAD_BYTES = 1 * 1024 * 1024 * 1024 # 1 GB max payload size 

1135 

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" 

1138 

1139 

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" 

1151 

1152 if not torch.cuda.is_available(): 

1153 logging.warning("CUDA not available, skipping distributed initialization.") 

1154 return 

1155 

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()) 

1167 

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) 

1174 

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 

1184 

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}") 

1189 

1190 

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 

1199 

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 ) 

1207 

1208 if dist.get_world_size() > 1: 

1209 init_distributed_environment( 

1210 rank=dist.get_rank(), 

1211 world_size=dist.get_world_size() 

1212 ) 

1213 

1214 if not dist.is_initialized(): 

1215 raise RuntimeError("Distributed process group not initialized.") 

1216 

1217 

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 

1226 

1227 logging.info(f"[{rank}] Waiting for all {local_world_size} workers to be ready...") 

1228 

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())}") 

1233 

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 

1239 

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}.") 

1243 

1244 logging.info(f"[{rank}] All {local_world_size} workers are ready.") 

1245 

1246 

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 

1264 

1265 global last_ping_time 

1266 last_ping_time = time.time() 

1267 

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") 

1273 

1274 if payload_size.item() > MAX_PAYLOAD_BYTES: 

1275 logging.error(f"[{rank}] Payload too large: {payload_size.item()} bytes.") 

1276 return 

1277 

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())}") 

1281 

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() 

1287 

1288 # Clean up the rank 0 signal file 

1289 await asyncio.to_thread(os.remove, SIGNAL_0_PATH) 

1290 

1291 dist.broadcast(payload_tensor, src=0) 

1292 torch.cuda.synchronize() 

1293 

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}") 

1299 

1300 

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 

1306 

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) 

1313 

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() 

1319 

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 

1324 

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() 

1329 

1330 payload_bytes = payload_tensor.cpu().numpy().tobytes() 

1331 payload = pickle.loads(payload_bytes) # nosec B301 - internal IPC from rank 0 

1332 

1333 if not isinstance(payload, dict): 

1334 logging.error(f"[{rank}] Invalid payload received: {payload}.") 

1335 continue 

1336 

1337 if not payload: 

1338 logging.error(f"[{rank}] Empty payload received.") 

1339 continue 

1340 

1341 if "task" not in payload: 

1342 logging.error(f"[{rank}] No 'task' in payload: {payload}.") 

1343 continue 

1344 

1345 gen_task = payload 

1346 

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) 

1374 

1375 await asyncio.sleep(0.1) # Breathing time between requests just in case 

1376 logging.info(f"[{rank}] Exiting worker thread.") 

1377 

1378 

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 

1385 

1386 

1387def get_model_names() -> List[str]: 

1388 """Get the list of available model names.""" 

1389 return list(models.keys()) 

1390 

1391 

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 

1399 

1400 

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") 

1407 

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") 

1412 

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") 

1452 

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 

1464 

1465 return args, engine_config 

1466 

1467 

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.") 

1476 

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 

1492 

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") 

1497 

1498 

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 ) 

1515 

1516 if args.wanvae: 

1517 model_name = "wanvae" 

1518 await load_model_wrapper_file(rank, model_name) 

1519 # TODO implement WanVideoVAEGeneration 

1520 

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) 

1530 

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) 

1536 

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) 

1542 

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() 

1548 

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) 

1554 

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) 

1560 

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) 

1566 

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) 

1572 

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) 

1578 

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) 

1584 

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) 

1590 

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) 

1596 

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) 

1602 

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) 

1608 

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) 

1614 

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) 

1620 

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) 

1626 

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() 

1632 

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() 

1638 

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() 

1644 

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() 

1650 

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() 

1656 

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) 

1662 

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) 

1668 

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() 

1674 

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() 

1680 

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() 

1686 

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() 

1692 

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() 

1698 

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() 

1704 

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() 

1710 

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() 

1716 

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() 

1722 

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() 

1728 

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.") 

1748 

1749 

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}"] 

1759 

1760 config.accesslog = "-" 

1761 

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 

1766 

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 

1771 

1772 if certfile: 

1773 config.certfile = certfile 

1774 if keyfile: 

1775 config.keyfile = keyfile 

1776 

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}") 

1781 

1782 logging.debug(f"[{rank}] HTTP server config:") 

1783 for key, value in app.config.items(): 

1784 logging.debug(f"{key}: {value}") 

1785 

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 

1792 

1793 

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() 

1803 

1804 args, engine_config = arg_parsing() 

1805 

1806 if not args.hunyuanavatar: 

1807 init_dist_environment() # Hunyuan-Avatar handles its own distributed environment 

1808 

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 )) 

1817 

1818 await init_model(args, engine_config) 

1819 

1820 await wait_for_everybody() 

1821 

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) 

1840 

1841 await wait_for_everybody() 

1842 

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.") 

1849 

1850 

1851if __name__ == "__main__": 

1852 setup_logging( 

1853 path=TMP_DIR, 

1854 file_name="streamwise.log", 

1855 level=logging.INFO) 

1856 

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)