|
18 | 18 | import logging |
19 | 19 | import operator |
20 | 20 | import os |
| 21 | +import tempfile |
21 | 22 | import time |
22 | 23 | import uuid |
23 | 24 | from concurrent.futures import ThreadPoolExecutor |
@@ -156,6 +157,12 @@ def load(self): |
156 | 157 | pipeline = self._model = HunyuanVideoPipeline.from_pretrained( |
157 | 158 | self._model_path, transformer=transformer, **kwargs |
158 | 159 | ) |
| 160 | + elif self.model_spec.model_family == "WanAnimate2": |
| 161 | + from diffusers import WanAnimate2Pipeline |
| 162 | + |
| 163 | + pipeline = self._model = WanAnimate2Pipeline.from_pretrained( |
| 164 | + self._model_path, **kwargs |
| 165 | + ) |
159 | 166 | elif self.model_spec.model_family == "Wan": |
160 | 167 | from diffusers import AutoencoderKLWan, WanImageToVideoPipeline, WanPipeline |
161 | 168 | from transformers import CLIPVisionModel |
@@ -312,9 +319,50 @@ def image_to_video( |
312 | 319 | assert callable(self._model) |
313 | 320 | generate_kwargs = self._model_spec.default_generate_config.copy() |
314 | 321 | generate_kwargs.update(kwargs) |
315 | | - generate_kwargs["num_videos_per_prompt"] = n |
316 | 322 | if num_inference_steps: |
317 | 323 | generate_kwargs["num_inference_steps"] = num_inference_steps |
| 324 | + |
| 325 | + if self.model_spec.model_family == "WanAnimate2": |
| 326 | + if n != 1: |
| 327 | + raise ValueError( |
| 328 | + "Wan-Animate-2 only supports generating one video per request" |
| 329 | + ) |
| 330 | + |
| 331 | + video = generate_kwargs.pop("video", None) |
| 332 | + if video is None: |
| 333 | + raise ValueError("`video` is required for Wan-Animate-2") |
| 334 | + |
| 335 | + fps = generate_kwargs.get("fps", 24) |
| 336 | + temp_video_path = None |
| 337 | + output: Any |
| 338 | + try: |
| 339 | + if isinstance(video, bytes): |
| 340 | + with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as f: |
| 341 | + temp_video_path = f.name |
| 342 | + f.write(video) |
| 343 | + video = temp_video_path |
| 344 | + elif isinstance(video, str): |
| 345 | + if not os.path.isfile(video): |
| 346 | + raise FileNotFoundError( |
| 347 | + f"Video path does not exist or is not a file: {video}" |
| 348 | + ) |
| 349 | + else: |
| 350 | + raise TypeError("`video` must be video bytes or a local path") |
| 351 | + |
| 352 | + self._process_progressor(generate_kwargs) |
| 353 | + output = self._model( |
| 354 | + image=image, |
| 355 | + driving_video=video, |
| 356 | + prompt=prompt, |
| 357 | + **generate_kwargs, |
| 358 | + ) |
| 359 | + finally: |
| 360 | + if temp_video_path and os.path.exists(temp_video_path): |
| 361 | + os.remove(temp_video_path) |
| 362 | + |
| 363 | + return self._output_to_video(output, fps, response_format) |
| 364 | + |
| 365 | + generate_kwargs["num_videos_per_prompt"] = n |
318 | 366 | fps = generate_kwargs.pop("fps", 10) |
319 | 367 |
|
320 | 368 | # process image |
|
0 commit comments