Instructions to use Kry4ta1/Effecteraser-VOR-Inference with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use Kry4ta1/Effecteraser-VOR-Inference with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("Kry4ta1/Effecteraser-VOR-Inference", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Download src/infer.py from Kry4ta1/Effecteraser-VOR-Inference: direct link, hf CLI and curl.
- Browser
- Download file 27.7 kB
-
https://huggingface.co/Kry4ta1/Effecteraser-VOR-Inference/resolve/main/src/infer.py
- Command line
-
hf download hf://Kry4ta1/Effecteraser-VOR-Inference/src/infer.py
-
curl -L -o infer.py https://huggingface.co/Kry4ta1/Effecteraser-VOR-Inference/resolve/main/src/infer.py
27.7 kB
| import os | |
| import argparse | |
| import numpy as np | |
| import torch | |
| import imageio | |
| from diffusers import FlowMatchEulerDiscreteScheduler | |
| from omegaconf import OmegaConf | |
| from PIL import Image | |
| from transformers import AutoTokenizer | |
| import scipy | |
| import cv2 | |
| from glob import glob | |
| import torch.distributed as dist | |
| from videox_fun.dist import set_multi_gpus_devices | |
| from videox_fun.models import AutoencoderKLWan, WanT5EncoderModel, VaceWanModel, load_lightx2v_vae | |
| from videox_fun.data.remove_dataset import orientation_aware_size | |
| from videox_fun.pipeline import RemovePipeline | |
| from videox_fun.utils.fp8_optimization import ( | |
| convert_model_weight_to_float8, | |
| replace_parameters_by_name, | |
| convert_weight_dtype_wrapper, | |
| ) | |
| from videox_fun.utils.lora_utils import merge_lora | |
| from videox_fun.utils.utils import save_videos_grid, filter_kwargs | |
| def load_patch_safetensors(path): | |
| list_tensors = glob(path + "/*.safetensors") | |
| all = {} | |
| for x in list_tensors: | |
| from safetensors.torch import load_file | |
| tmp = load_file(x) | |
| all.update(tmp) | |
| return all | |
| def parse_args(): | |
| parser = argparse.ArgumentParser(description="WanFun Video Editing Script") | |
| # GPU and memory configuration | |
| parser.add_argument( | |
| "--gpu_memory_mode", | |
| type=str, | |
| default="model_full_load", | |
| choices=["model_full_load", "model_cpu_offload", "model_cpu_offload_and_qfloat8", "sequential_cpu_offload"], | |
| help="GPU memory optimization mode", | |
| ) | |
| parser.add_argument("--ulysses_degree", type=int, default=1, help="Ulysses degree for multi-GPU configuration") | |
| parser.add_argument("--ring_degree", type=int, default=1, help="Ring degree for multi-GPU configuration") | |
| # Model paths | |
| parser.add_argument( | |
| "--config_path", type=str, default="config/wan2.1/wan_civitai.yaml", help="Path to model configuration file" | |
| ) | |
| parser.add_argument("--model_name", type=str, default="models/Wan2.1-VACE-1.3B", help="Path to pretrained model") | |
| parser.add_argument( | |
| "--common_model_name", | |
| type=str, | |
| default=None, | |
| help="Optional directory containing the shared VAE, text encoder, and tokenizer. Defaults to --model_name.", | |
| ) | |
| # VAE configuration. The Remove LoRAs are trained against the pruned LightVAE latent | |
| # distribution, so inference must use the same LightVAE by default. Using the full | |
| # Wan VAE both mismatches the latents and blows up decode-time GPU memory (OOM). | |
| parser.add_argument( | |
| "--lightvae_path", | |
| type=str, | |
| default="/data1/yfu/1_CVPR2025_Remove/10_transsion/6_0713/lightvaew2_1.pth", | |
| help="Path to the LightVAE checkpoint used during training. Loaded by default.", | |
| ) | |
| parser.add_argument("--lightvae_pruning_rate", type=float, default=0.75, help="LightVAE channel pruning rate") | |
| parser.add_argument("--lightvae_dim", type=int, default=96, help="LightVAE base model dim") | |
| parser.add_argument( | |
| "--use_full_vae", | |
| action="store_true", | |
| help="Force the full official Wan VAE instead of LightVAE (higher VRAM; only for debugging).", | |
| ) | |
| # Generation parameters | |
| parser.add_argument( | |
| "--sample_size", | |
| type=str, | |
| default="480,832", | |
| help="Orientation-aware canvas as 'height,width' for landscape; swapped for portrait. Matches training.", | |
| ) | |
| parser.add_argument("--video_length", type=int, default=81, help="Length of generated video in frames") | |
| parser.add_argument("--fps", type=int, default=16, help="Frames per second for output video") | |
| parser.add_argument( | |
| "--weight_dtype", | |
| type=str, | |
| default="bfloat16", | |
| choices=["float16", "bfloat16"], | |
| help="Data type for model weights", | |
| ) | |
| # Prompt and generation settings | |
| parser.add_argument( | |
| "--prompt", | |
| type=str, | |
| default="Remove the target and fill the content appropriately", | |
| help="Text prompt for generation", | |
| ) | |
| parser.add_argument( | |
| "--negative_prompt", | |
| type=str, | |
| default="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走", | |
| help="Negative text prompt", | |
| ) | |
| parser.add_argument( | |
| "--guidance_scale", | |
| type=float, | |
| default=1.0, | |
| help="Guidance scale. 1.0 disables CFG (single forward per step); >1.0 enables CFG (2x forwards).", | |
| ) | |
| parser.add_argument("--seed", type=int, default=43, help="Random seed for reproducibility") | |
| parser.add_argument("--context_scale", type=float, default=1.0, help="Context scale for vace control") | |
| parser.add_argument("--dilation", type=int, default=6, help="Dilation for inp mask (only for inpaint mode)") | |
| # Parameters for Remove | |
| parser.add_argument( | |
| "--lora_path", | |
| type=str, | |
| default=["models/remove_model_stage1.safetensors", "models/remove_model_stage2.safetensors"], | |
| nargs="+", | |
| help="Optional path to LoRA checkpoint", | |
| ) | |
| parser.add_argument( | |
| "--lora_weight", type=float, default=[1.0, 1.0], nargs="+", help="Weight for LoRA model if used" | |
| ) | |
| parser.add_argument( | |
| "--skip_lora", | |
| action="store_true", | |
| help="Do not load external LoRAs. Required when --model_name already contains merged LoRAs.", | |
| ) | |
| # Single-pair inference | |
| parser.add_argument( | |
| "--input_video", | |
| type=str, | |
| default=None, | |
| help="Path to a single input video for editing", | |
| ) | |
| parser.add_argument( | |
| "--input_mask_video", | |
| type=str, | |
| default=None, | |
| help="Path to a single mask video for editing", | |
| ) | |
| # Directory batch inference | |
| parser.add_argument( | |
| "--input_dir", | |
| type=str, | |
| default=None, | |
| help="Directory containing input videos", | |
| ) | |
| parser.add_argument( | |
| "--input_mask_dir", | |
| type=str, | |
| default=None, | |
| help="Directory containing mask videos with exactly matching filenames", | |
| ) | |
| parser.add_argument("--num_inference_steps", type=int, default=4, help="Number of inference steps") | |
| parser.add_argument("--dmd_steps", type=int, choices=[1, 2], default=None) | |
| parser.add_argument("--shard_index", type=int, default=0) | |
| parser.add_argument("--num_shards", type=int, default=1) | |
| parser.add_argument("--save_dir", type=str, default="samples/Remove", help="Directory to save generated videos") | |
| return parser.parse_args() | |
| def process_video( | |
| input_video_path, | |
| input_mask_video_path, | |
| video_length, | |
| sample_size, | |
| dilation=0, | |
| ): | |
| """Process input video and mask for editing""" | |
| if input_video_path is not None: | |
| cap = cv2.VideoCapture(input_video_path) | |
| frames = [] | |
| while cap.isOpened(): | |
| ret, frame = cap.read() | |
| if not ret: | |
| break | |
| frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) | |
| frames.append(Image.fromarray(frame)) | |
| cap.release() | |
| frames = frames[:video_length] | |
| if len(frames) < video_length: | |
| frames += [frames[-1]] * (video_length - len(frames)) | |
| resized_frames = [frame.resize([sample_size[1], sample_size[0]]) for frame in frames] | |
| # Keep the raw RGB frames (uint8, [T,H,W,3]) for the cat/overlay visualization. | |
| raw_frames = np.stack([np.array(frame) for frame in resized_frames]).astype(np.uint8) | |
| input_video = ( | |
| torch.stack([torch.from_numpy(np.array(frame)).permute(2, 0, 1) for frame in resized_frames]) | |
| .permute(1, 0, 2, 3) | |
| .unsqueeze(0) | |
| ) # [1, C, T, H, W] | |
| else: | |
| input_video = torch.zeros((1, 3, video_length, sample_size[0], sample_size[1])).float() | |
| raw_frames = np.zeros((video_length, sample_size[0], sample_size[1], 3), dtype=np.uint8) | |
| if input_mask_video_path is not None: | |
| mask_cap = cv2.VideoCapture(input_mask_video_path) | |
| mask_frames = [] | |
| while mask_cap.isOpened(): | |
| ret, frame = mask_cap.read() | |
| if not ret: | |
| break | |
| if len(frame.shape) == 3: | |
| frame = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY) | |
| _, mask = cv2.threshold(frame, 127, 255, cv2.THRESH_BINARY) | |
| if dilation > 0: | |
| mask_np = (mask > 0).astype(np.uint8) | |
| mask = scipy.ndimage.binary_dilation(mask_np, iterations=dilation).astype(np.uint8) * 255 | |
| mask_frames.append(mask) | |
| mask_cap.release() | |
| mask_frames = mask_frames[:video_length] | |
| if len(mask_frames) < video_length: | |
| mask_frames += [mask_frames[-1]] * (video_length - len(mask_frames)) | |
| resized_masks = [Image.fromarray(mask).resize([sample_size[1], sample_size[0]]) for mask in mask_frames] | |
| # Keep the binary mask (uint8, [T,H,W], 0/255) for the cat/overlay visualization. | |
| raw_masks = np.stack([np.array(mask) for mask in resized_masks]).astype(np.uint8) | |
| input_video_mask = ( | |
| torch.stack([torch.from_numpy(np.array(mask)) for mask in resized_masks]).unsqueeze(0).unsqueeze(0) / 255.0 | |
| ) # [1, 1, T, H, W] | |
| else: | |
| input_video_mask = torch.ones((1, 1, video_length, sample_size[0], sample_size[1])).float() | |
| raw_masks = np.full((video_length, sample_size[0], sample_size[1]), 255, dtype=np.uint8) | |
| if input_video_path is not None and input_video is not None: | |
| input_video = input_video * (torch.tile(input_video_mask, [1, 3, 1, 1, 1]) < 0.5) + (128.0) * ( | |
| torch.tile(input_video_mask, [1, 3, 1, 1, 1]) >= 0.5 | |
| ) | |
| input_video = input_video.div_(127.5).sub_(1.0) | |
| return input_video, input_video_mask, raw_frames, raw_masks | |
| def process_single_task( | |
| pipeline, | |
| args, | |
| input_video_path, | |
| input_mask_video_path, | |
| prompt, | |
| ): | |
| """Process a single video editing task""" | |
| if input_video_path is not None: | |
| # Get video resolution | |
| cap = cv2.VideoCapture(input_video_path) | |
| width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)) | |
| height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) | |
| cap.release() | |
| # Match training: resize to a fixed orientation-aware canvas rather than an | |
| # aspect-ratio-preserving max-area box. Training always feeds landscape | |
| # (height,width) or its swap for portrait, so inference must do the same or | |
| # the model sees an out-of-distribution resolution. | |
| base_height, base_width = [int(value) for value in args.sample_size.split(",")] | |
| new_height, new_width = orientation_aware_size(width, height, base_height, base_width) | |
| sample_size = [new_height, new_width] | |
| else: | |
| sample_size = [int(args.sample_size.split(",")[0]), int(args.sample_size.split(",")[1])] | |
| generator = torch.Generator(device=pipeline.device).manual_seed(args.seed) | |
| with torch.no_grad(): | |
| video_length = ( | |
| int( | |
| (args.video_length - 1) | |
| // pipeline.vae.config.temporal_compression_ratio | |
| * pipeline.vae.config.temporal_compression_ratio | |
| ) | |
| + 1 | |
| if args.video_length != 1 | |
| else 1 | |
| ) | |
| # Process video and mask | |
| ( | |
| input_video, | |
| input_video_mask, | |
| raw_frames, | |
| raw_masks, | |
| ) = process_video( | |
| input_video_path, | |
| input_mask_video_path, | |
| video_length=video_length, | |
| sample_size=sample_size, | |
| dilation=args.dilation, | |
| ) | |
| # Generate edited video | |
| sample = pipeline( | |
| prompt, | |
| negative_prompt=args.negative_prompt, | |
| height=sample_size[0], | |
| width=sample_size[1], | |
| generator=generator, | |
| guidance_scale=args.guidance_scale, | |
| num_inference_steps=args.num_inference_steps, | |
| video=input_video, | |
| mask_video=input_video_mask, | |
| context_scale=args.context_scale, | |
| ).videos | |
| if not torch.isfinite(sample).all(): | |
| raise FloatingPointError('Inference generated NaN or infinite pixel values') | |
| timing = pipeline.last_timing | |
| print( | |
| "[Timing] " | |
| f"vae_encode={timing.get('vae_encode_seconds', float('nan')):.3f}s, " | |
| f"condition_prepare={timing.get('condition_prepare_seconds', float('nan')):.3f}s, " | |
| f"denoise={timing.get('denoise_seconds', float('nan')):.3f}s, " | |
| f"vae_decode={timing.get('vae_decode_seconds', float('nan')):.3f}s, " | |
| f"decode_postprocess={timing.get('decode_postprocess_seconds', float('nan')):.3f}s, " | |
| f"vae_total={timing.get('vae_total_seconds', float('nan')):.3f}s, " | |
| f"pipeline_total={timing.get('pipeline_total_seconds', float('nan')):.3f}s" | |
| ) | |
| return sample, video_length, raw_frames, raw_masks | |
| def sample_to_uint8_frames(sample): | |
| """Convert a pipeline sample [1,3,T,H,W] in [0,1] to uint8 RGB frames [T,H,W,3].""" | |
| video = sample[0].detach().float().clamp(0, 1).cpu() # [3,T,H,W] | |
| frames = video.permute(1, 2, 3, 0).mul(255).round().clamp(0, 255).to(torch.uint8).numpy() | |
| return frames # [T,H,W,3] | |
| def overlay_mask_on_frames( | |
| frames, | |
| masks, | |
| fill_color=(255, 255, 0), | |
| edge_color=(255, 0, 0), | |
| fill_alpha=0.3, | |
| edge_thickness=2, | |
| ): | |
| """Annotate the masked region on each frame. | |
| The region is filled with a highly transparent yellow (low alpha) and outlined | |
| with an opaque red edge, so the target area is clearly marked without hiding | |
| the underlying content. | |
| Args: | |
| frames: uint8 RGB array [T,H,W,3]. | |
| masks: uint8 array [T,H,W] with 0/255 (or any >0 as foreground). | |
| fill_color: RGB of the transparent fill (default yellow). | |
| edge_color: RGB of the opaque outline (default red). | |
| fill_alpha: opacity of the fill in [0,1]; small means highly transparent. | |
| edge_thickness: outline thickness in pixels. | |
| """ | |
| fill = np.array(fill_color, dtype=np.float32) | |
| edge = np.array(edge_color, dtype=np.uint8) | |
| erode_kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (3, 3)) | |
| out = [] | |
| for frame, mask in zip(frames, masks): | |
| frame = frame.astype(np.float32) | |
| binary = (mask > 127).astype(np.uint8) | |
| if binary.any(): | |
| # Highly transparent yellow fill inside the region. | |
| selected = binary.astype(bool) | |
| frame[selected] = (1.0 - fill_alpha) * frame[selected] + fill_alpha * fill | |
| # Opaque red outline along the region boundary. | |
| eroded = cv2.erode(binary, erode_kernel, iterations=edge_thickness) | |
| border = (binary - eroded).astype(bool) | |
| frame_uint8 = frame.clip(0, 255).astype(np.uint8) | |
| frame_uint8[border] = edge | |
| else: | |
| frame_uint8 = frame.clip(0, 255).astype(np.uint8) | |
| out.append(frame_uint8) | |
| return np.stack(out) | |
| def save_video_frames(frames, path, fps): | |
| """Write uint8 RGB frames [T,H,W,3] to an mp4 via imageio.""" | |
| os.makedirs(os.path.dirname(path), exist_ok=True) | |
| imageio.mimsave(path, list(frames), fps=fps) | |
| def save_results(sample, args, video_length, fps, task_name=None, raw_frames=None, raw_masks=None): | |
| """Save the generated results. | |
| Produces two artifacts: | |
| - meta/<name>.mp4 : the generated result on its own. | |
| - cat/<name>.mp4 : the input video with the mask region annotated (highly | |
| transparent yellow fill + red edge), horizontally concatenated with the | |
| generated result. | |
| """ | |
| prefix = task_name | |
| meta_dir = os.path.join(args.save_dir, "meta") | |
| cat_dir = os.path.join(args.save_dir, "cat") | |
| os.makedirs(meta_dir, exist_ok=True) | |
| if video_length == 1: | |
| # Single-frame (image) output. | |
| meta_path = os.path.join(meta_dir, prefix + ".png") | |
| image = sample[0, :, 0] | |
| image = image.transpose(0, 1).transpose(1, 2) | |
| gen_image = (image * 255).numpy().astype(np.uint8) | |
| Image.fromarray(gen_image).save(meta_path) | |
| if raw_frames is not None and raw_masks is not None: | |
| os.makedirs(cat_dir, exist_ok=True) | |
| overlay = overlay_mask_on_frames(raw_frames[:1], raw_masks[:1])[0] | |
| cat_image = np.concatenate([overlay, gen_image], axis=1) | |
| Image.fromarray(cat_image).save(os.path.join(cat_dir, prefix + ".png")) | |
| return | |
| # Video output. | |
| gen_frames = sample_to_uint8_frames(sample) # [T,H,W,3] | |
| meta_path = os.path.join(meta_dir, prefix + ".mp4") | |
| save_video_frames(gen_frames, meta_path, fps) | |
| if raw_frames is not None and raw_masks is not None: | |
| os.makedirs(cat_dir, exist_ok=True) | |
| frame_count = min(len(raw_frames), len(raw_masks), len(gen_frames)) | |
| overlay = overlay_mask_on_frames(raw_frames[:frame_count], raw_masks[:frame_count]) | |
| cat_frames = np.concatenate([overlay, gen_frames[:frame_count]], axis=2) # concat along width | |
| save_video_frames(cat_frames, os.path.join(cat_dir, prefix + ".mp4"), fps) | |
| def main(): | |
| args = parse_args() | |
| if not 0 <= args.shard_index < args.num_shards: | |
| raise ValueError('Require 0 <= shard_index < num_shards') | |
| if args.dmd_steps is not None: | |
| if args.guidance_scale != 1.0: | |
| raise ValueError('DMD student is distilled at CFG=1; inference must use CFG=1') | |
| args.num_inference_steps = args.dmd_steps | |
| # Validate arguments | |
| single_mode = args.input_video is not None or args.input_mask_video is not None | |
| dir_mode = args.input_dir is not None or args.input_mask_dir is not None | |
| if single_mode and dir_mode: | |
| raise ValueError( | |
| "Do not mix single-file mode and directory mode. " | |
| "Use either --input_video + --input_mask_video OR --input_dir + --input_mask_dir." | |
| ) | |
| if single_mode: | |
| if args.input_video is None or args.input_mask_video is None: | |
| raise ValueError("Single-file mode requires both --input_video and --input_mask_video") | |
| if not os.path.isfile(args.input_video): | |
| raise FileNotFoundError(f"Input video not found: {args.input_video}") | |
| if not os.path.isfile(args.input_mask_video): | |
| raise FileNotFoundError(f"Input mask video not found: {args.input_mask_video}") | |
| elif dir_mode: | |
| if args.input_dir is None or args.input_mask_dir is None: | |
| raise ValueError("Directory mode requires both --input_dir and --input_mask_dir") | |
| if not os.path.isdir(args.input_dir): | |
| raise NotADirectoryError(f"Input directory not found: {args.input_dir}") | |
| if not os.path.isdir(args.input_mask_dir): | |
| raise NotADirectoryError(f"Mask directory not found: {args.input_mask_dir}") | |
| else: | |
| raise ValueError( | |
| "Please provide either --input_video + --input_mask_video " | |
| "or --input_dir + --input_mask_dir" | |
| ) | |
| # Convert weight dtype | |
| weight_dtype = torch.bfloat16 if args.weight_dtype == "bfloat16" else torch.float16 | |
| device = set_multi_gpus_devices(args.ulysses_degree, args.ring_degree) | |
| config = OmegaConf.load(args.config_path) | |
| common_model_name = args.common_model_name or args.model_name | |
| # Initialize transformer | |
| from load_vace import load_vace, load_text_encoder | |
| transformer = load_vace( | |
| os.path.join( | |
| args.model_name, config["transformer_additional_kwargs"].get("transformer_subpath", "transformer") | |
| ), | |
| OmegaConf.to_container(config["transformer_additional_kwargs"]), | |
| weight_dtype, | |
| ) | |
| # Get Vae. Default to the pruned LightVAE the Remove LoRAs were trained with; only | |
| # fall back to the full Wan VAE when explicitly requested via --use_full_vae. | |
| if args.use_full_vae: | |
| print("[INFO] Using full official Wan VAE (--use_full_vae). This needs more GPU memory.") | |
| vae = AutoencoderKLWan.from_pretrained( | |
| os.path.join(common_model_name, config["vae_kwargs"].get("vae_subpath", "vae")), | |
| additional_kwargs=OmegaConf.to_container(config["vae_kwargs"]), | |
| ).to(weight_dtype) | |
| else: | |
| if not args.lightvae_path or not os.path.isfile(args.lightvae_path): | |
| raise FileNotFoundError( | |
| f"LightVAE checkpoint not found: {args.lightvae_path}. " | |
| "Pass a valid --lightvae_path, or use --use_full_vae to force the full VAE." | |
| ) | |
| print( | |
| f"[INFO] Loading LightVAE from {args.lightvae_path} " | |
| f"(pruning_rate={args.lightvae_pruning_rate}, dim={args.lightvae_dim})" | |
| ) | |
| vae = load_lightx2v_vae( | |
| args.lightvae_path, | |
| pruning_rate=args.lightvae_pruning_rate, | |
| model_dim=args.lightvae_dim, | |
| torch_dtype=weight_dtype, | |
| device="cpu", | |
| ) | |
| # Get Tokenizer | |
| tokenizer = AutoTokenizer.from_pretrained( | |
| os.path.join(common_model_name, config["text_encoder_kwargs"].get("tokenizer_subpath", "tokenizer")), | |
| ) | |
| # Get Text encoder | |
| text_encoder = load_text_encoder( | |
| os.path.join(common_model_name, config["text_encoder_kwargs"].get("text_encoder_subpath", "text_encoder")), | |
| OmegaConf.to_container(config["text_encoder_kwargs"]),weight_dtype, | |
| ) | |
| text_encoder = text_encoder.eval() | |
| # Get Scheduler | |
| Choosen_Scheduler = { | |
| "Flow": FlowMatchEulerDiscreteScheduler, | |
| }["Flow"] | |
| scheduler = Choosen_Scheduler( | |
| **filter_kwargs(Choosen_Scheduler, OmegaConf.to_container(config["scheduler_kwargs"])) | |
| ) | |
| if args.dmd_steps is not None: | |
| from dmd_scheduler import DMDFlowScheduler | |
| scheduler = DMDFlowScheduler.from_config(scheduler.config) | |
| scheduler.dmd_steps = args.dmd_steps | |
| # Get Pipeline | |
| pipeline = RemovePipeline( | |
| transformer=transformer, | |
| vae=vae, | |
| tokenizer=tokenizer, | |
| text_encoder=text_encoder, | |
| scheduler=scheduler, | |
| ) | |
| if args.ulysses_degree > 1 or args.ring_degree > 1: | |
| transformer.enable_multi_gpus_inference() | |
| if args.gpu_memory_mode == "sequential_cpu_offload": | |
| replace_parameters_by_name( | |
| transformer, | |
| [ | |
| "modulation", | |
| ], | |
| device=device, | |
| ) | |
| transformer.freqs = transformer.freqs.to(device=device) | |
| pipeline.enable_sequential_cpu_offload(device=device) | |
| elif args.gpu_memory_mode == "model_cpu_offload_and_qfloat8": | |
| convert_model_weight_to_float8( | |
| transformer, | |
| exclude_module_name=[ | |
| "modulation", | |
| ], | |
| ) | |
| convert_weight_dtype_wrapper(transformer, weight_dtype) | |
| pipeline.enable_model_cpu_offload(device=device) | |
| elif args.gpu_memory_mode == "model_cpu_offload": | |
| pipeline.enable_model_cpu_offload(device=device) | |
| else: | |
| pipeline.to(device=device) | |
| if args.skip_lora: | |
| print("[INFO] Skipping external LoRA loading; using transformer weights from --model_name directly.") | |
| elif args.lora_path is not None: | |
| if len(args.lora_weight) != len(args.lora_path): | |
| args.lora_weight = [args.lora_weight[0]] * len(args.lora_path) | |
| for lora_path, lora_weight in zip(args.lora_path, args.lora_weight): | |
| print(f"[INFO] Loading LoRA: {lora_path}, weight: {lora_weight}") | |
| pipeline = merge_lora(pipeline, lora_path, lora_weight) | |
| if not os.path.exists(args.save_dir): | |
| os.makedirs(args.save_dir, exist_ok=True) | |
| # Inference | |
| if args.input_dir is not None: | |
| # Batch directory inference. Model/pipeline is loaded only once above. | |
| valid_exts = {".mp4", ".avi", ".mov", ".mkv", ".webm"} | |
| input_filenames = sorted( | |
| filename | |
| for filename in os.listdir(args.input_dir) | |
| if os.path.isfile(os.path.join(args.input_dir, filename)) | |
| and os.path.splitext(filename)[1].lower() in valid_exts | |
| ) | |
| pairs = [] | |
| for filename in input_filenames: | |
| input_video_path = os.path.join(args.input_dir, filename) | |
| input_mask_video_path = os.path.join(args.input_mask_dir, filename) | |
| # Exact filename matching: abc.mp4 <-> abc.mp4 | |
| if os.path.isfile(input_mask_video_path): | |
| pairs.append((filename, input_video_path, input_mask_video_path)) | |
| else: | |
| if not dist.is_initialized() or dist.get_rank() == 0: | |
| print(f"[WARNING] No matching mask for {filename}; skipping.") | |
| if len(pairs) == 0: | |
| raise RuntimeError( | |
| f"No valid video-mask pairs found. " | |
| f"input_dir={args.input_dir}, input_mask_dir={args.input_mask_dir}" | |
| ) | |
| pairs = pairs[args.shard_index::args.num_shards] | |
| if not dist.is_initialized() or dist.get_rank() == 0: | |
| print(f"[INFO] Found {len(pairs)} valid video-mask pairs.") | |
| for idx, (filename, input_video_path, input_mask_video_path) in enumerate(pairs, start=1): | |
| if not dist.is_initialized() or dist.get_rank() == 0: | |
| print("\n" + "=" * 80) | |
| print(f"[INFO] Processing {idx}/{len(pairs)}: {filename}") | |
| print(f"[INFO] Video: {input_video_path}") | |
| print(f"[INFO] Mask : {input_mask_video_path}") | |
| print("=" * 80) | |
| sample, video_length, raw_frames, raw_masks = process_single_task( | |
| pipeline, | |
| args, | |
| input_video_path, | |
| input_mask_video_path, | |
| args.prompt, | |
| ) | |
| video_basename = os.path.splitext(filename)[0] | |
| if not dist.is_initialized() or dist.get_rank() == 0: | |
| save_results( | |
| sample, | |
| args, | |
| video_length, | |
| args.fps, | |
| task_name=video_basename, | |
| raw_frames=raw_frames, | |
| raw_masks=raw_masks, | |
| ) | |
| print( | |
| f"[INFO] Saved: meta={os.path.join(args.save_dir, 'meta', video_basename + '.mp4')}, " | |
| f"cat={os.path.join(args.save_dir, 'cat', video_basename + '.mp4')}" | |
| ) | |
| # Release per-sample tensors before the next pair. | |
| del sample | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| if not dist.is_initialized() or dist.get_rank() == 0: | |
| print(f"\n[INFO] Finished all {len(pairs)} pairs. Results: {args.save_dir}") | |
| else: | |
| # Original single-pair inference | |
| sample, video_length, raw_frames, raw_masks = process_single_task( | |
| pipeline, | |
| args, | |
| args.input_video, | |
| args.input_mask_video, | |
| args.prompt, | |
| ) | |
| video_basename = os.path.splitext(os.path.basename(args.input_video))[0] | |
| if not dist.is_initialized() or dist.get_rank() == 0: | |
| save_results( | |
| sample, | |
| args, | |
| video_length, | |
| args.fps, | |
| task_name=video_basename, | |
| raw_frames=raw_frames, | |
| raw_masks=raw_masks, | |
| ) | |
| if __name__ == "__main__": | |
| main() | |