Coverage for wrapper/hunyuanframepack/wrapper_hunyuanframepack_base.py: 64%

384 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-09 04:47 +0000

1import logging 

2import math 

3import time 

4import tempfile 

5import aiofiles 

6 

7import numpy as np 

8 

9from typing import Union 

10from typing import Optional 

11from typing import Dict 

12from typing import Any 

13from typing import Tuple 

14 

15import torch 

16import torch.distributed as dist 

17 

18from torch import inference_mode 

19 

20from model_timing import GenTimer 

21from wrapper_usp import USPGeneration 

22from media_utils import save_bcthw_as_mp4 

23 

24from PIL import Image 

25 

26from diffusers import AutoencoderKLHunyuanVideo 

27from diffusers import HunyuanVideoFramepackPipeline 

28from diffusers import FlowMatchEulerDiscreteScheduler 

29 

30from diffusers_helper.hunyuan import vae_decode 

31from diffusers_helper.hunyuan import vae_encode 

32from diffusers_helper.hunyuan import encode_prompt_conds 

33from diffusers_helper.utils import crop_or_pad_yield_mask 

34from diffusers_helper.utils import repeat_to_batch_size 

35from diffusers_helper.utils import resize_and_center_crop 

36from diffusers_helper.clip_vision import hf_clip_vision_encode 

37from diffusers_helper.models.hunyuan_video_packed import HunyuanVideoTransformer3DModelPacked 

38from diffusers_helper.pipelines.k_diffusion_hunyuan import get_flux_sigmas_from_mu 

39from diffusers_helper.k_diffusion.uni_pc_fm import FlowMatchUniPC 

40from diffusers_helper.k_diffusion.wrapper import fm_wrapper 

41 

42from transformers import LlamaModel 

43from transformers import CLIPTextModel 

44from transformers import LlamaTokenizerFast 

45from transformers import CLIPTokenizer 

46from transformers import SiglipImageProcessor 

47from transformers import SiglipVisionModel 

48 

49from image_utils import base64_to_img 

50 

51if torch.cuda.is_available(): 

52 from hunyuanframepack_xfuser import parallelize_transformer 

53 

54from xfuser.config import EngineConfig 

55from xfuser.core.distributed import initialize_runtime_state 

56from xfuser.core.distributed import get_runtime_state 

57 

58 

59def get_hidden_size( 

60 height: int, 

61 width: int, 

62) -> int: 

63 lat_h = height / 8 

64 lat_w = width / 8 

65 lat_h_pad4 = (lat_h + 3) // 4 * 4 

66 lat_w_pad4 = (lat_w + 3) // 4 * 4 

67 lat_h_pad8 = (lat_h + 7) // 8 * 8 

68 lat_w_pad8 = (lat_w + 7) // 8 * 8 

69 dim = int(9 * lat_h / 2 * lat_w / 2 + lat_h * lat_w / 2 + lat_h_pad4 

70 / 4 * lat_w_pad4 / 4 + 4 * lat_h_pad8 / 8 * lat_w_pad8 / 8) 

71 return dim 

72 

73 

74class HunyuanFramePackBase(USPGeneration): 

75 def __init__( 

76 self, 

77 model_name: str = "hunyuanframepack", 

78 framepack_model_name: str = "lllyasviel/FramePackI2V_HY", 

79 engine_config: Optional[EngineConfig] = None, 

80 param_dtype: torch.dtype = torch.bfloat16, 

81 enable_tiling: bool = False, 

82 enable_slicing: bool = False, 

83 ) -> None: 

84 super().__init__( 

85 model_name=model_name, 

86 engine_config=engine_config, 

87 param_dtype=param_dtype, 

88 ) 

89 

90 self.enable_tiling = enable_tiling 

91 self.enable_slicing = enable_slicing 

92 

93 # Model components 

94 self.text_encoder: Optional[LlamaModel] = None 

95 self.text_encoder_2: Optional[CLIPTextModel] = None 

96 self.tokenizer: Optional[LlamaTokenizerFast] = None 

97 self.tokenizer_2: Optional[CLIPTokenizer] = None 

98 self.vae: Optional[AutoencoderKLHunyuanVideo] = None 

99 self.scheduler: Optional[FlowMatchEulerDiscreteScheduler] = None 

100 self.feature_extractor: Optional[SiglipImageProcessor] = None 

101 self.image_encoder: Optional[SiglipVisionModel] = None 

102 self.transformer: Optional[HunyuanVideoTransformer3DModelPacked] = None 

103 

104 # Model features 

105 self.framepack_model_name = framepack_model_name 

106 self.shift = 3.0 

107 self.strength = 1.0 

108 self.vae_stride = (4, 8, 8) # time, height, width 

109 self.LAT_CHANNELS = 16 

110 self.num_heads = 24 

111 self.FPS = 30 # This is technically a constant for the model 

112 

113 def __del__(self) -> None: 

114 # Clean models 

115 if self.text_encoder is not None: 

116 del self.text_encoder 

117 if self.text_encoder_2 is not None: 

118 del self.text_encoder_2 

119 if self.tokenizer is not None: 

120 del self.tokenizer 

121 if self.tokenizer_2 is not None: 

122 del self.tokenizer_2 

123 if self.vae is not None: 

124 del self.vae 

125 if self.scheduler is not None: 

126 del self.scheduler 

127 if self.feature_extractor is not None: 

128 del self.feature_extractor 

129 if self.image_encoder is not None: 

130 del self.image_encoder 

131 if self.transformer is not None: 

132 del self.transformer 

133 super().__del__() 

134 

135 def load_model(self) -> None: 

136 assert torch.cuda.is_available() 

137 assert self.device is not None 

138 

139 prev_memory = torch.cuda.memory_allocated() 

140 self.load_timer.start("text_encoder") 

141 self.text_encoder = LlamaModel.from_pretrained( 

142 "hunyuanvideo-community/HunyuanVideo", 

143 subfolder="text_encoder", 

144 torch_dtype=torch.float16).to(self.device) # nosec B615 

145 assert self.text_encoder is not None 

146 self.text_encoder.eval().requires_grad_(False) 

147 self.load_timer.end("text_encoder") 

148 diff_memory = torch.cuda.memory_allocated() - prev_memory 

149 logging.info(f"[{self.rank}] Text memory allocated: {diff_memory / 1024 / 1024 ** 2:.2f} GB.") 

150 

151 self.text_encoder_2 = CLIPTextModel.from_pretrained( 

152 "hunyuanvideo-community/HunyuanVideo", 

153 subfolder='text_encoder_2', 

154 torch_dtype=torch.float16).to(self.device) # nosec B615 

155 assert self.text_encoder_2 is not None 

156 self.text_encoder_2.eval().requires_grad_(False) 

157 

158 self.tokenizer = LlamaTokenizerFast.from_pretrained( 

159 "hunyuanvideo-community/HunyuanVideo/tokenizer") # nosec B615 

160 

161 self.tokenizer_2 = CLIPTokenizer.from_pretrained( 

162 "hunyuanvideo-community/HunyuanVideo", 

163 subfolder='tokenizer_2') # nosec B615 

164 

165 prev_memory = torch.cuda.memory_allocated() if torch.cuda.is_available() else 0 

166 self.load_timer.start("vae") 

167 self.vae = AutoencoderKLHunyuanVideo.from_pretrained( 

168 "hunyuanvideo-community/HunyuanVideo", 

169 subfolder='vae', 

170 torch_dtype=torch.float16).to(self.device) 

171 assert self.vae is not None 

172 self.vae.eval().requires_grad_(False) # type: ignore[union-attr] 

173 

174 if not self.enable_tiling: 

175 logging.info(f"[{self.rank}] Disabling tiling for VAE.") 

176 self.vae.disable_tiling() # type: ignore[union-attr] 

177 else: 

178 logging.info(f"[{self.rank}] Enabling tiling for VAE.") 

179 self.vae.enable_tiling() # type: ignore[union-attr] 

180 

181 if not self.enable_slicing: 

182 logging.info(f"[{self.rank}] Disabling slicing for VAE.") 

183 self.vae.disable_slicing() # type: ignore[union-attr] 

184 else: 

185 logging.info(f"[{self.rank}] Enabling slicing for VAE.") 

186 self.vae.enable_slicing() # type: ignore[union-attr] 

187 self.load_timer.end("vae") 

188 

189 diff_memory = torch.cuda.memory_allocated() - prev_memory 

190 logging.info(f"[{self.rank}] VAE memory allocated: {diff_memory / 1024 / 1024 ** 2:.2f} GB.") 

191 

192 self.scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( 

193 "hunyuanvideo-community/HunyuanVideo", 

194 subfolder='scheduler', 

195 torch_dtype=torch.float16) 

196 

197 self.feature_extractor = SiglipImageProcessor.from_pretrained( 

198 "lllyasviel/flux_redux_bfl", 

199 subfolder='feature_extractor') # nosec B615 

200 

201 prev_memory = torch.cuda.memory_allocated() 

202 self.load_timer.start("image_encoder") 

203 self.image_encoder = SiglipVisionModel.from_pretrained( 

204 "lllyasviel/flux_redux_bfl", 

205 subfolder='image_encoder', 

206 torch_dtype=torch.float16).to(self.device) # nosec B615 

207 self.image_encoder.eval().requires_grad_(False) 

208 self.load_timer.end("image_encoder") 

209 diff_memory = torch.cuda.memory_allocated() - prev_memory 

210 logging.info(f"[{self.rank}] Img memory allocated: {diff_memory / 1024 / 1024 ** 2:.2f} GB.") 

211 

212 prev_memory = torch.cuda.memory_allocated() 

213 self.load_timer.start("dit") 

214 self.transformer = HunyuanVideoTransformer3DModelPacked.from_pretrained( 

215 self.framepack_model_name, 

216 torch_dtype=self.param_dtype 

217 ) 

218 self.transformer = self.transformer.to(self.device) 

219 assert self.transformer is not None 

220 self.transformer.eval().requires_grad_(False) 

221 self.load_timer.end("dit") 

222 diff_memory = torch.cuda.memory_allocated() - prev_memory 

223 logging.info(f"[{self.rank}] FramePack memory allocated: {diff_memory / 1024 / 1024 ** 2:.2f} GB.") 

224 

225 def init_model_parallelism(self) -> None: 

226 if not dist.is_initialized() or self.world_size <= 1: 

227 return 

228 

229 self.load_timer.start("dit_parallel") 

230 assert self.text_encoder is not None 

231 assert self.tokenizer is not None 

232 assert self.transformer is not None 

233 assert self.vae is not None 

234 assert self.scheduler is not None 

235 assert self.text_encoder_2 is not None 

236 assert self.tokenizer_2 is not None 

237 assert self.image_encoder is not None 

238 assert self.feature_extractor is not None 

239 temp_pipeline = HunyuanVideoFramepackPipeline( 

240 self.text_encoder, 

241 self.tokenizer, 

242 self.transformer, 

243 self.vae, 

244 self.scheduler, 

245 self.text_encoder_2, 

246 self.tokenizer_2, 

247 self.image_encoder, 

248 self.feature_extractor, 

249 ) 

250 assert temp_pipeline is not None 

251 initialize_runtime_state(temp_pipeline, self.engine_config) 

252 get_runtime_state().set_video_input_parameters( 

253 batch_size=1, 

254 ) 

255 parallelize_transformer(temp_pipeline) 

256 self.load_timer.end("dit_parallel") 

257 

258 def model_compile(self) -> None: 

259 if not self.torch_compile: 

260 return 

261 

262 # This started to happen with torch 2.8.0 

263 # Skipping dit_compile as it fails with: KeyError: op23 Set TORCHDYNAMO_VERBOSE=1" 

264 logging.info(f"[{self.rank}] Compiling DiT with torch.compile().") 

265 self.load_timer.start("dit_compile") 

266 self.transformer = torch.compile( 

267 self.transformer, 

268 mode="max-autotune-no-cudagraphs", 

269 ) 

270 self.load_timer.end("dit_compile") 

271 

272 logging.info(f"[{self.rank}] Compiling VAE with torch.compile().") 

273 self.load_timer.start("vae_compile") 

274 assert self.vae is not None 

275 self.vae = torch.compile( # type: ignore[call-overload] 

276 self.vae, 

277 mode="max-autotune-no-cudagraphs", 

278 ) 

279 self.load_timer.end("vae_compile") 

280 

281 def _assert_model_init(self) -> None: 

282 super()._assert_model_init() 

283 assert self.text_encoder is not None 

284 assert self.image_encoder is not None 

285 assert self.vae is not None 

286 assert self.transformer is not None 

287 

288 def _assert_args( 

289 self, 

290 height: int, 

291 width: int, 

292 ) -> None: 

293 """Check if the image size is supported for the current parallelism.""" 

294 if height % self.vae_stride[1] != 0: 

295 raise ValueError(f"Height {height} must be divisible by {self.vae_stride[1]}.") 

296 if width % self.vae_stride[2] != 0: 

297 raise ValueError(f"Width {width} must be divisible by {self.vae_stride[2]}.") 

298 hidden_size = get_hidden_size(height, width) 

299 if hidden_size % self.world_size != 0: 

300 raise ValueError(f"{height}x{width} ({hidden_size}) not supported for {self.world_size} GPUs.") 

301 

302 @inference_mode() 

303 async def warmup(self) -> None: 

304 logging.info(f"[{self.rank}] Warmup for Hunyuan FramePack ({self.model_name}) generation.") 

305 await self.generate( 

306 job_id="warmup", 

307 img=Image.new("RGB", (768, 512), (255, 255, 255)), 

308 prompt="Warmup prompt", 

309 neg_prompt="", 

310 height=512, 

311 width=768, 

312 num_frames=1 + 4, 

313 sampling_steps=5) 

314 

315 def _encode_text( 

316 self, 

317 gen_timer: GenTimer, 

318 prompt: str, 

319 neg_prompt: str, 

320 cfg: float, 

321 ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: 

322 """ Text encoder for Hunyuan FramePack generation.""" 

323 gen_timer.start("text_encoder") 

324 llama_vec, clip_l_pooler = encode_prompt_conds( 

325 prompt, 

326 self.text_encoder, 

327 self.text_encoder_2, 

328 self.tokenizer, 

329 self.tokenizer_2) 

330 if cfg == 1: 

331 llama_vec_n = torch.zeros_like(llama_vec) 

332 clip_l_pooler_n = torch.zeros_like(clip_l_pooler) 

333 else: 

334 llama_vec_n, clip_l_pooler_n = encode_prompt_conds( 

335 neg_prompt, 

336 self.text_encoder, 

337 self.text_encoder_2, 

338 self.tokenizer, 

339 self.tokenizer_2) 

340 llama_vec, llama_attention_mask = crop_or_pad_yield_mask(llama_vec, length=512) 

341 llama_vec_n, llama_attention_mask_n = crop_or_pad_yield_mask(llama_vec_n, length=512) 

342 assert self.transformer is not None 

343 llama_vec = llama_vec.to(self.transformer.dtype) 

344 llama_vec_n = llama_vec_n.to(self.transformer.dtype) 

345 clip_l_pooler = clip_l_pooler.to(self.transformer.dtype) 

346 clip_l_pooler_n = clip_l_pooler_n.to(self.transformer.dtype) 

347 gen_timer.end("text_encoder") 

348 

349 return llama_vec, llama_attention_mask, clip_l_pooler, llama_vec_n, llama_attention_mask_n, clip_l_pooler_n 

350 

351 def _process_image( 

352 self, 

353 gen_timer: GenTimer, 

354 img: Image.Image, 

355 height: int, 

356 width: int, 

357 ) -> Tuple[np.ndarray, torch.Tensor]: 

358 """ Process image for Hunyuan FramePack generation.""" 

359 # h,w,RGB -> 1,RGB,1,h,w ([544,704,3] -> [1,3,1,544,704]) 

360 gen_timer.start("image_preprocess") 

361 # We may resize image for parallelism; 1024x720 works with SP=8 

362 img_resized = img.resize((width, height), Image.Resampling.LANCZOS) 

363 input_image = np.array(img_resized) 

364 t0_img = time.time() 

365 H, W, C = input_image.shape 

366 if C != 3: 

367 raise ValueError(f"Input image must be RGB: {input_image.shape}") 

368 # The model works with other resolutions, skip buckets 

369 # height, width = find_nearest_bucket(H, W, resolution=640) # 720x1280 -> 480x832 

370 input_image_np = resize_and_center_crop(input_image, target_width=width, target_height=height) 

371 input_image_pt = torch.from_numpy(input_image_np).float() / (255.0 / 2.0) - 1 

372 input_image_pt = input_image_pt.permute(2, 0, 1)[None, :, None] 

373 gen_timer.end("image_preprocess") 

374 if self.rank == 0: 

375 logging.info( 

376 f"[{self.rank}] Image processing time: {time.time() - t0_img:.3f} seconds " 

377 f"img:{W}x{H}, video:{width}x{height}.") 

378 return input_image_np, input_image_pt 

379 

380 def _clip_vision( 

381 self, 

382 gen_timer: GenTimer, 

383 input_image_np: np.ndarray, 

384 ) -> torch.Tensor: 

385 """ CLIP Vision encoder for Hunyuan FramePack generation.""" 

386 # (1, 729, 1152) 

387 gen_timer.start("image_encoder") 

388 t0_clip = time.time() 

389 image_encoder_output = hf_clip_vision_encode(input_image_np, self.feature_extractor, self.image_encoder) 

390 image_encoder_last_hidden_state = image_encoder_output.last_hidden_state 

391 assert self.transformer is not None 

392 image_encoder_last_hidden_state = image_encoder_last_hidden_state.to(self.transformer.dtype) 

393 gen_timer.end("image_encoder") 

394 if self.rank == 0: 

395 logging.info(f"[{self.rank}] CLIP Vision encoding time: {time.time() - t0_clip:.3f} seconds.") 

396 return image_encoder_last_hidden_state 

397 

398 @inference_mode() 

399 async def generate( 

400 self, 

401 img: Image.Image, 

402 prompt: str, 

403 neg_prompt: str = "", 

404 height: int = 512, 

405 width: int = 768, 

406 num_frames: int = 1 + 80, 

407 sampling_steps: int = 25, # 10 

408 # latent frames for every Hunyuan Video window: 9->36 pixel frames -> 1.2 seconds 

409 latent_window_size: int = 9, 

410 cfg: float = 1.0, 

411 distilled_guidance_scale: float = 10.0, 

412 guidance_rescale: int = 0, 

413 save_intermediate: Optional[str] = None, 

414 job_id: Optional[str] = None, 

415 output_type: str = "tensor" 

416 ) -> torch.Tensor: 

417 raise NotImplementedError("Implement generate() in subclasses.") 

418 

419 @inference_mode() 

420 def vae_decode( 

421 self, 

422 latents: torch.Tensor, 

423 job_id: Optional[str] = None, 

424 ) -> torch.Tensor: 

425 """ 

426 Latent -> Pixels 

427 """ 

428 gen_timer = self._new_gen_timer(job_id) 

429 

430 assert self.vae is not None 

431 assert latents is not None 

432 assert isinstance(latents, torch.Tensor) 

433 if latents.ndim != 5: # B, C, T, H, W 

434 raise ValueError(f"Latents must be a 5D tensor (B, C, T, H, W), got {latents.ndim}D.") 

435 if latents.shape[1] != 16: 

436 raise ValueError(f"Latents must have 16 channels, got {latents.shape[1]} channels.") 

437 

438 try: 

439 gen_timer.start("vae_decoder") 

440 latents = latents.to(self.device, dtype=self.param_dtype) 

441 pixels = vae_decode(latents, self.vae) 

442 gen_timer.end("vae_decoder") 

443 return pixels 

444 finally: 

445 gen_timer.end("total") 

446 

447 @inference_mode() 

448 def vae_encode( 

449 self, 

450 pixels: torch.Tensor, 

451 job_id: Optional[str] = None, 

452 ) -> torch.Tensor: 

453 """ 

454 Pixels -> Latent 

455 """ 

456 gen_timer = self._new_gen_timer(job_id) 

457 

458 assert self.vae is not None 

459 assert pixels is not None 

460 assert isinstance(pixels, torch.Tensor) 

461 if pixels.ndim != 5: 

462 raise ValueError(f"Pixels must be a 5D tensor (B, C, T, H, W), got {pixels.ndim}D.") 

463 if pixels.shape[1] != 3: 

464 raise ValueError(f"Pixels must have 3 channels (RGB), got {pixels.shape[1]} channels.") 

465 

466 try: 

467 gen_timer.start("vae_encoder") 

468 pixels = pixels.to(self.device, dtype=self.param_dtype) 

469 latents = vae_encode(pixels, self.vae) 

470 gen_timer.end("vae_encoder") 

471 return latents 

472 finally: 

473 gen_timer.end("total") 

474 

475 @inference_mode() 

476 def _sample_hunyuan( 

477 self, 

478 it0: int, 

479 gen_timer: GenTimer, 

480 initial_latent: Optional[torch.Tensor] = None, 

481 concat_latent: Optional[torch.Tensor] = None, 

482 strength: float = 1.0, 

483 width: int = 512, 

484 height: int = 512, 

485 frames: int = 16, 

486 real_guidance_scale: float = 1.0, 

487 distilled_guidance_scale: float = 6.0, 

488 guidance_rescale: float = 0.0, 

489 num_inference_steps: int = 25, 

490 batch_size: Optional[int] = None, 

491 generator: Optional[Any] = None, 

492 prompt_embeds: Optional[torch.Tensor] = None, 

493 prompt_embeds_mask: Optional[torch.Tensor] = None, 

494 prompt_poolers: Optional[torch.Tensor] = None, 

495 negative_prompt_embeds: Optional[torch.Tensor] = None, 

496 negative_prompt_embeds_mask: Optional[torch.Tensor] = None, 

497 negative_prompt_poolers: Optional[torch.Tensor] = None, 

498 negative_kwargs: Optional[Dict[str, Any]] = None, 

499 **kwargs: Any, 

500 ) -> torch.Tensor: 

501 if batch_size is None: 

502 assert prompt_embeds is not None 

503 batch_size = int(prompt_embeds.shape[0]) 

504 

505 # Random noise 

506 # B, C, T, H, W (1, 16, X, 8, 8) 

507 LAT_CHANNELS = 16 

508 assert generator is not None 

509 assert self.vae_stride is not None 

510 latents = torch.randn( 

511 ( 

512 batch_size, 

513 LAT_CHANNELS, 

514 (frames + 3) // self.vae_stride[0], 

515 height // self.vae_stride[1], 

516 width // self.vae_stride[2] 

517 ), 

518 generator=generator, device=generator.device 

519 ).to(device=self.device, dtype=torch.float32) 

520 

521 mu = math.log(self.shift) 

522 

523 sigmas = get_flux_sigmas_from_mu(num_inference_steps, mu).to(self.device) 

524 

525 k_model = fm_wrapper(self.transformer) 

526 

527 if initial_latent is not None: 

528 sigmas = sigmas * strength 

529 first_sigma = sigmas[0].to(device=self.device, dtype=torch.float32) 

530 initial_latent = initial_latent.to(device=self.device, dtype=torch.float32) 

531 latents = initial_latent.float() * (1.0 - first_sigma) + latents.float() * first_sigma 

532 

533 if concat_latent is not None: 

534 concat_latent = concat_latent.to(latents) 

535 

536 distilled_guidance = torch.tensor([distilled_guidance_scale * 1000.0] 

537 * batch_size).to(device=self.device, dtype=self.param_dtype) 

538 

539 prompt_embeds = repeat_to_batch_size(prompt_embeds, batch_size) 

540 prompt_embeds_mask = repeat_to_batch_size(prompt_embeds_mask, batch_size) 

541 prompt_poolers = repeat_to_batch_size(prompt_poolers, batch_size) 

542 negative_prompt_embeds = repeat_to_batch_size(negative_prompt_embeds, batch_size) 

543 negative_prompt_embeds_mask = repeat_to_batch_size(negative_prompt_embeds_mask, batch_size) 

544 negative_prompt_poolers = repeat_to_batch_size(negative_prompt_poolers, batch_size) 

545 concat_latent = repeat_to_batch_size(concat_latent, batch_size) 

546 

547 sampler_kwargs = dict( 

548 dtype=self.param_dtype, 

549 cfg_scale=real_guidance_scale, 

550 cfg_rescale=guidance_rescale, 

551 concat_latent=concat_latent, 

552 positive=dict( 

553 pooled_projections=prompt_poolers, 

554 encoder_hidden_states=prompt_embeds, 

555 encoder_attention_mask=prompt_embeds_mask, 

556 guidance=distilled_guidance, 

557 **kwargs, 

558 ), 

559 negative=dict( 

560 pooled_projections=negative_prompt_poolers, 

561 encoder_hidden_states=negative_prompt_embeds, 

562 encoder_attention_mask=negative_prompt_embeds_mask, 

563 guidance=distilled_guidance, 

564 **(kwargs if negative_kwargs is None else {**kwargs, **negative_kwargs}), 

565 ) 

566 ) 

567 

568 sampler = FlowMatchUniPC(k_model, extra_args=sampler_kwargs) 

569 

570 # generated_latents = sampler.sample(latents, sigmas=sigmas) 

571 order = min(3, len(sigmas) - 2) 

572 model_prev_list, t_prev_list = [], [] 

573 for it1 in range(len(sigmas) - 1): 

574 gen_timer.start(f"dit_{it0:03d}_{it1:03d}") 

575 vec_t = sigmas[it1].expand(latents.shape[0]) 

576 if it1 == 0: 

577 model_prev_list = [sampler.model_fn(latents, vec_t)] 

578 t_prev_list = [vec_t] 

579 elif it1 < order: 

580 init_order = it1 

581 latents, model_x = sampler.update_fn(latents, model_prev_list, t_prev_list, vec_t, init_order) 

582 model_prev_list.append(model_x) 

583 t_prev_list.append(vec_t) 

584 else: 

585 latents, model_x = sampler.update_fn(latents, model_prev_list, t_prev_list, vec_t, order) 

586 model_prev_list.append(model_x) 

587 t_prev_list.append(vec_t) 

588 model_prev_list = model_prev_list[-order:] 

589 t_prev_list = t_prev_list[-order:] 

590 gen_timer.end(f"dit_{it0:03d}_{it1:03d}") 

591 generated_latents = model_prev_list[-1] 

592 

593 return generated_latents 

594 

595 async def _output_video( 

596 self, 

597 job_id: Optional[str], 

598 gen_timer: GenTimer, 

599 pixels: torch.Tensor, 

600 output_type: str = "tensor", # "tensor", "video_binary", "video_path" 

601 ) -> Union[torch.Tensor, str, bytes, None]: 

602 gen_timer.start("output") 

603 try: 

604 if output_type == "tensor": 

605 return pixels 

606 

607 if output_type in ("video_binary", "video_path"): 

608 if not job_id: 

609 video_path = tempfile.NamedTemporaryFile(suffix=".mp4", delete=False).name 

610 else: 

611 video_path = f"/tmp/{job_id}.mp4" 

612 video_path = save_bcthw_as_mp4( 

613 pixels, 

614 video_path, 

615 fps=self.FPS) 

616 if output_type == "video_path": 

617 return video_path 

618 

619 # video_binary 

620 async with aiofiles.open(video_path, "rb") as file: 

621 video_binary = await file.read() 

622 return video_binary 

623 

624 logging.error(f"Unknown output type: {output_type}") 

625 return None 

626 finally: 

627 gen_timer.end("output") 

628 

629 async def get_rest_args( 

630 self, 

631 data_json: Dict[str, Union[str, int, float]] 

632 ) -> Dict[str, Any]: 

633 if data_json is None or not isinstance(data_json, dict): 

634 raise ValueError("Missing JSON body") 

635 

636 job_id = data_json.get("job_id", None) 

637 

638 img_base64 = data_json.get("img", None) 

639 if img_base64 is None: 

640 raise ValueError("Missing 'img' parameter") 

641 if not isinstance(img_base64, str): 

642 raise ValueError("'img' parameter must be a base64 string") 

643 img = base64_to_img(img_base64) 

644 

645 prompt = data_json.get("prompt", None) 

646 if prompt is None: 

647 raise ValueError("Missing 'prompt' parameter") 

648 

649 neg_prompt = data_json.get("neg_prompt", "") 

650 

651 height = int(data_json.get("height", 480)) 

652 width = int(data_json.get("width", 640)) 

653 num_frames = int(data_json.get("num_frames", 1 + 16)) 

654 steps = int(data_json.get("sampling_steps", 5)) 

655 latent_window_size = int(data_json.get("latent_window_size", 9)) 

656 cfg = float(data_json.get("cfg", 1.0)) 

657 distilled_guidance_scale = float(data_json.get("distilled_guidance_scale", 10.0)) 

658 guidance_rescale = int(data_json.get("guidance_rescale", 0.0)) 

659 save_intermediate = data_json.get("save_intermediate", None) 

660 output_type = data_json.get("output_type", "tensor") 

661 

662 if height <= 0: 

663 raise ValueError(f"height {height} must be positive.") 

664 if width <= 0: 

665 raise ValueError(f"width {width} must be positive.") 

666 if steps <= 0: 

667 raise ValueError(f"sampling_steps {steps} must be positive.") 

668 if latent_window_size <= 0: 

669 raise ValueError(f"latent_window_size {latent_window_size} must be positive.") 

670 

671 video_seconds = data_json.get("video_seconds", None) 

672 if video_seconds is not None: 

673 if float(video_seconds) <= 0: 

674 raise ValueError(f"video_seconds {video_seconds} must be positive.") 

675 VAE_FRAMES = self.vae_stride[0] 

676 num_frames = int(video_seconds * self.FPS) 

677 num_frames = 1 + ((num_frames - 1) // VAE_FRAMES) * VAE_FRAMES # 4n + 1 

678 

679 if num_frames <= 0: 

680 raise ValueError(f"num_frames {num_frames} must be positive.") 

681 

682 return { 

683 "task": self.model_name, 

684 "args": { 

685 "job_id": job_id, 

686 "img": img, 

687 "prompt": prompt, 

688 "neg_prompt": neg_prompt, 

689 "height": height, 

690 "width": width, 

691 "num_frames": num_frames, 

692 "sampling_steps": steps, 

693 "latent_window_size": latent_window_size, 

694 "cfg": cfg, 

695 "distilled_guidance_scale": distilled_guidance_scale, 

696 "guidance_rescale": guidance_rescale, 

697 "save_intermediate": save_intermediate, 

698 "output_type": output_type, 

699 } 

700 }