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
« 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
7import numpy as np
9from typing import Union
10from typing import Optional
11from typing import Dict
12from typing import Any
13from typing import Tuple
15import torch
16import torch.distributed as dist
18from torch import inference_mode
20from model_timing import GenTimer
21from wrapper_usp import USPGeneration
22from media_utils import save_bcthw_as_mp4
24from PIL import Image
26from diffusers import AutoencoderKLHunyuanVideo
27from diffusers import HunyuanVideoFramepackPipeline
28from diffusers import FlowMatchEulerDiscreteScheduler
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
42from transformers import LlamaModel
43from transformers import CLIPTextModel
44from transformers import LlamaTokenizerFast
45from transformers import CLIPTokenizer
46from transformers import SiglipImageProcessor
47from transformers import SiglipVisionModel
49from image_utils import base64_to_img
51if torch.cuda.is_available():
52 from hunyuanframepack_xfuser import parallelize_transformer
54from xfuser.config import EngineConfig
55from xfuser.core.distributed import initialize_runtime_state
56from xfuser.core.distributed import get_runtime_state
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
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 )
90 self.enable_tiling = enable_tiling
91 self.enable_slicing = enable_slicing
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
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
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__()
135 def load_model(self) -> None:
136 assert torch.cuda.is_available()
137 assert self.device is not None
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.")
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)
158 self.tokenizer = LlamaTokenizerFast.from_pretrained(
159 "hunyuanvideo-community/HunyuanVideo/tokenizer") # nosec B615
161 self.tokenizer_2 = CLIPTokenizer.from_pretrained(
162 "hunyuanvideo-community/HunyuanVideo",
163 subfolder='tokenizer_2') # nosec B615
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]
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]
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")
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.")
192 self.scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained(
193 "hunyuanvideo-community/HunyuanVideo",
194 subfolder='scheduler',
195 torch_dtype=torch.float16)
197 self.feature_extractor = SiglipImageProcessor.from_pretrained(
198 "lllyasviel/flux_redux_bfl",
199 subfolder='feature_extractor') # nosec B615
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.")
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.")
225 def init_model_parallelism(self) -> None:
226 if not dist.is_initialized() or self.world_size <= 1:
227 return
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")
258 def model_compile(self) -> None:
259 if not self.torch_compile:
260 return
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")
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")
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
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.")
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)
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")
349 return llama_vec, llama_attention_mask, clip_l_pooler, llama_vec_n, llama_attention_mask_n, clip_l_pooler_n
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
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
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.")
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)
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.")
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")
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)
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.")
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")
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])
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)
521 mu = math.log(self.shift)
523 sigmas = get_flux_sigmas_from_mu(num_inference_steps, mu).to(self.device)
525 k_model = fm_wrapper(self.transformer)
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
533 if concat_latent is not None:
534 concat_latent = concat_latent.to(latents)
536 distilled_guidance = torch.tensor([distilled_guidance_scale * 1000.0]
537 * batch_size).to(device=self.device, dtype=self.param_dtype)
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)
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 )
568 sampler = FlowMatchUniPC(k_model, extra_args=sampler_kwargs)
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]
593 return generated_latents
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
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
619 # video_binary
620 async with aiofiles.open(video_path, "rb") as file:
621 video_binary = await file.read()
622 return video_binary
624 logging.error(f"Unknown output type: {output_type}")
625 return None
626 finally:
627 gen_timer.end("output")
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")
636 job_id = data_json.get("job_id", None)
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)
645 prompt = data_json.get("prompt", None)
646 if prompt is None:
647 raise ValueError("Missing 'prompt' parameter")
649 neg_prompt = data_json.get("neg_prompt", "")
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")
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.")
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
679 if num_frames <= 0:
680 raise ValueError(f"num_frames {num_frames} must be positive.")
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 }