Download inference/source/src/fastwam/runtime.py from Recharge23/FastWAM-single: direct link, hf CLI and curl.
- Browser
- Download file 21.9 kB
-
https://huggingface.co/Recharge23/FastWAM-single/resolve/main/inference/source/src/fastwam/runtime.py
- Command line
-
hf download hf://Recharge23/FastWAM-single/inference/source/src/fastwam/runtime.py
-
curl -L -o runtime.py https://huggingface.co/Recharge23/FastWAM-single/resolve/main/inference/source/src/fastwam/runtime.py
21.9 kB
| import logging | |
| import os | |
| import inspect | |
| from pathlib import Path | |
| import torch | |
| from hydra.utils import instantiate | |
| from omegaconf import DictConfig | |
| from PIL import Image | |
| import numpy as np | |
| from einops import repeat | |
| from omegaconf import OmegaConf | |
| from .utils.logging_config import get_logger, setup_logging | |
| from .utils.video_io import save_mp4 | |
| from .utils import misc | |
| logger = get_logger(__name__) | |
| def _normalize_mixed_precision(mixed_precision: str) -> str: | |
| if not isinstance(mixed_precision, str): | |
| raise ValueError(f"`mixed_precision` must be str, got {type(mixed_precision)}") | |
| key = mixed_precision.strip().lower() | |
| if key not in {"no", "fp16", "bf16"}: | |
| raise ValueError( | |
| f"Unsupported mixed_precision: {mixed_precision}. " | |
| "Expected one of: ['no', 'fp16', 'bf16']." | |
| ) | |
| return key | |
| def _mixed_precision_to_model_dtype(mixed_precision: str) -> torch.dtype: | |
| precision = _normalize_mixed_precision(mixed_precision) | |
| if precision == "no": | |
| return torch.float32 | |
| if precision == "fp16": | |
| return torch.float16 | |
| return torch.bfloat16 | |
| def create_wan22_model( | |
| model_id: str, | |
| tokenizer_model_id: str, | |
| dit_config, | |
| tokenizer_max_len: int = 512, | |
| train_shift: float = 5.0, | |
| infer_shift: float = 5.0, | |
| num_train_timesteps: int = 1000, | |
| redirect_common_files: bool = True, | |
| model_dtype: torch.dtype = torch.bfloat16, | |
| device: str = "cuda", | |
| ): | |
| from .models.wan22.wan22 import Wan22Core | |
| if isinstance(dit_config, DictConfig): | |
| dit_config = OmegaConf.to_container(dit_config, resolve=True) | |
| if not isinstance(dit_config, dict): | |
| raise ValueError(f"`dit_config` must resolve to a dict, got {type(dit_config)}") | |
| return Wan22Core.from_wan22_pretrained( | |
| device=device, | |
| torch_dtype=model_dtype, | |
| model_id=model_id, | |
| tokenizer_model_id=tokenizer_model_id, | |
| tokenizer_max_len=int(tokenizer_max_len), | |
| redirect_common_files=bool(redirect_common_files), | |
| dit_config=dit_config, | |
| train_shift=float(train_shift), | |
| infer_shift=float(infer_shift), | |
| num_train_timesteps=int(num_train_timesteps), | |
| ) | |
| def create_fastwam( | |
| model_id: str, | |
| tokenizer_model_id: str, | |
| video_dit_config, | |
| tokenizer_max_len: int = 512, | |
| load_text_encoder: bool = True, | |
| proprio_dim: int | None = None, | |
| action_dit_config=None, | |
| action_dit_pretrained_path: str | None = None, | |
| skip_dit_load_from_pretrain: bool = False, | |
| video_scheduler=None, | |
| action_scheduler=None, | |
| loss=None, | |
| mot_checkpoint_mixed_attn: bool = False, | |
| compile_training_denoise: bool = False, | |
| redirect_common_files: bool = True, | |
| model_dtype: torch.dtype = torch.bfloat16, | |
| device: str = "cuda", | |
| ): | |
| from .models.wan22.fastwam import FastWAM | |
| if isinstance(video_dit_config, DictConfig): | |
| video_dit_config = OmegaConf.to_container(video_dit_config, resolve=True) | |
| if not isinstance(video_dit_config, dict): | |
| raise ValueError(f"`video_dit_config` must resolve to a dict, got {type(video_dit_config)}") | |
| if isinstance(action_dit_config, DictConfig): | |
| action_dit_config = OmegaConf.to_container(action_dit_config, resolve=True) | |
| if action_dit_config is None: | |
| action_dit_config = {} | |
| if not isinstance(action_dit_config, dict): | |
| raise ValueError(f"`action_dit_config` must resolve to a dict, got {type(action_dit_config)}") | |
| if isinstance(video_scheduler, DictConfig): | |
| video_scheduler = OmegaConf.to_container(video_scheduler, resolve=True) | |
| if video_scheduler is None: | |
| video_scheduler = {} | |
| if not isinstance(video_scheduler, dict): | |
| raise ValueError(f"`video_scheduler` must be dict-like, got {type(video_scheduler)}") | |
| if isinstance(action_scheduler, DictConfig): | |
| action_scheduler = OmegaConf.to_container(action_scheduler, resolve=True) | |
| if action_scheduler is None: | |
| raise ValueError("`action_scheduler` is required for FastWAM.") | |
| if not isinstance(action_scheduler, dict): | |
| raise ValueError(f"`action_scheduler` must be dict-like, got {type(action_scheduler)}") | |
| required_action_scheduler_keys = {"train_shift", "infer_shift", "num_train_timesteps"} | |
| missing_keys = required_action_scheduler_keys - set(action_scheduler.keys()) | |
| if missing_keys: | |
| raise ValueError( | |
| f"`action_scheduler` missing required keys: {sorted(missing_keys)}. " | |
| "Expected keys: train_shift, infer_shift, num_train_timesteps." | |
| ) | |
| if isinstance(loss, DictConfig): | |
| loss = OmegaConf.to_container(loss, resolve=True) | |
| if loss is None: | |
| loss = {} | |
| if not isinstance(loss, dict): | |
| raise ValueError(f"`loss` must be dict-like, got {type(loss)}") | |
| return FastWAM.from_wan22_pretrained( | |
| device=device, | |
| torch_dtype=model_dtype, | |
| model_id=model_id, | |
| tokenizer_model_id=tokenizer_model_id, | |
| tokenizer_max_len=int(tokenizer_max_len), | |
| load_text_encoder=bool(load_text_encoder), | |
| proprio_dim=(None if proprio_dim is None else int(proprio_dim)), | |
| redirect_common_files=bool(redirect_common_files), | |
| video_dit_config=video_dit_config, | |
| action_dit_config=action_dit_config, | |
| action_dit_pretrained_path=action_dit_pretrained_path, | |
| skip_dit_load_from_pretrain=bool(skip_dit_load_from_pretrain), | |
| mot_checkpoint_mixed_attn=bool(mot_checkpoint_mixed_attn), | |
| video_train_shift=float(video_scheduler.get("train_shift", 5.0)), | |
| video_infer_shift=float(video_scheduler.get("infer_shift", 5.0)), | |
| video_num_train_timesteps=int(video_scheduler.get("num_train_timesteps", 1000)), | |
| action_train_shift=float(action_scheduler["train_shift"]), | |
| action_infer_shift=float(action_scheduler["infer_shift"]), | |
| action_num_train_timesteps=int(action_scheduler["num_train_timesteps"]), | |
| loss_lambda_video=float(loss.get("lambda_video", 1.0)), | |
| loss_lambda_action=float(loss.get("lambda_action", 1.0)), | |
| compile_training_denoise=bool(compile_training_denoise), | |
| ) | |
| def create_fastwam_joint( | |
| model_id: str, | |
| tokenizer_model_id: str, | |
| video_dit_config, | |
| tokenizer_max_len: int = 512, | |
| load_text_encoder: bool = True, | |
| proprio_dim: int | None = None, | |
| action_dit_config=None, | |
| action_dit_pretrained_path: str | None = None, | |
| skip_dit_load_from_pretrain: bool = False, | |
| video_scheduler=None, | |
| action_scheduler=None, | |
| loss=None, | |
| mot_checkpoint_mixed_attn: bool = False, | |
| redirect_common_files: bool = True, | |
| model_dtype: torch.dtype = torch.bfloat16, | |
| device: str = "cuda", | |
| ): | |
| from .models.wan22.fastwam_joint import FastWAMJoint | |
| if isinstance(video_dit_config, DictConfig): | |
| video_dit_config = OmegaConf.to_container(video_dit_config, resolve=True) | |
| if not isinstance(video_dit_config, dict): | |
| raise ValueError(f"`video_dit_config` must resolve to a dict, got {type(video_dit_config)}") | |
| if isinstance(action_dit_config, DictConfig): | |
| action_dit_config = OmegaConf.to_container(action_dit_config, resolve=True) | |
| if action_dit_config is None: | |
| action_dit_config = {} | |
| if not isinstance(action_dit_config, dict): | |
| raise ValueError(f"`action_dit_config` must resolve to a dict, got {type(action_dit_config)}") | |
| if isinstance(video_scheduler, DictConfig): | |
| video_scheduler = OmegaConf.to_container(video_scheduler, resolve=True) | |
| if video_scheduler is None: | |
| video_scheduler = {} | |
| if not isinstance(video_scheduler, dict): | |
| raise ValueError(f"`video_scheduler` must be dict-like, got {type(video_scheduler)}") | |
| if isinstance(action_scheduler, DictConfig): | |
| action_scheduler = OmegaConf.to_container(action_scheduler, resolve=True) | |
| if action_scheduler is None: | |
| raise ValueError("`action_scheduler` is required for FastWAM.") | |
| if not isinstance(action_scheduler, dict): | |
| raise ValueError(f"`action_scheduler` must be dict-like, got {type(action_scheduler)}") | |
| required_action_scheduler_keys = {"train_shift", "infer_shift", "num_train_timesteps"} | |
| missing_keys = required_action_scheduler_keys - set(action_scheduler.keys()) | |
| if missing_keys: | |
| raise ValueError( | |
| f"`action_scheduler` missing required keys: {sorted(missing_keys)}. " | |
| "Expected keys: train_shift, infer_shift, num_train_timesteps." | |
| ) | |
| if isinstance(loss, DictConfig): | |
| loss = OmegaConf.to_container(loss, resolve=True) | |
| if loss is None: | |
| loss = {} | |
| if not isinstance(loss, dict): | |
| raise ValueError(f"`loss` must be dict-like, got {type(loss)}") | |
| return FastWAMJoint.from_wan22_pretrained( | |
| device=device, | |
| torch_dtype=model_dtype, | |
| model_id=model_id, | |
| tokenizer_model_id=tokenizer_model_id, | |
| tokenizer_max_len=int(tokenizer_max_len), | |
| load_text_encoder=bool(load_text_encoder), | |
| proprio_dim=(None if proprio_dim is None else int(proprio_dim)), | |
| redirect_common_files=bool(redirect_common_files), | |
| video_dit_config=video_dit_config, | |
| action_dit_config=action_dit_config, | |
| action_dit_pretrained_path=action_dit_pretrained_path, | |
| skip_dit_load_from_pretrain=bool(skip_dit_load_from_pretrain), | |
| mot_checkpoint_mixed_attn=bool(mot_checkpoint_mixed_attn), | |
| video_train_shift=float(video_scheduler.get("train_shift", 5.0)), | |
| video_infer_shift=float(video_scheduler.get("infer_shift", 5.0)), | |
| video_num_train_timesteps=int(video_scheduler.get("num_train_timesteps", 1000)), | |
| action_train_shift=float(action_scheduler["train_shift"]), | |
| action_infer_shift=float(action_scheduler["infer_shift"]), | |
| action_num_train_timesteps=int(action_scheduler["num_train_timesteps"]), | |
| loss_lambda_video=float(loss.get("lambda_video", 1.0)), | |
| loss_lambda_action=float(loss.get("lambda_action", 1.0)), | |
| ) | |
| def create_fastwam_idm( | |
| model_id: str, | |
| tokenizer_model_id: str, | |
| video_dit_config, | |
| video_cond_noise_prob: float = 0.5, | |
| tokenizer_max_len: int = 512, | |
| load_text_encoder: bool = True, | |
| proprio_dim: int | None = None, | |
| action_dit_config=None, | |
| action_dit_pretrained_path: str | None = None, | |
| skip_dit_load_from_pretrain: bool = False, | |
| video_scheduler=None, | |
| action_scheduler=None, | |
| loss=None, | |
| mot_checkpoint_mixed_attn: bool = False, | |
| compile_training_denoise: bool = False, | |
| redirect_common_files: bool = True, | |
| model_dtype: torch.dtype = torch.bfloat16, | |
| device: str = "cuda", | |
| ): | |
| from .models.wan22.fastwam_idm import ( | |
| FastWAMIDM, | |
| ) | |
| if isinstance(video_dit_config, DictConfig): | |
| video_dit_config = OmegaConf.to_container(video_dit_config, resolve=True) | |
| if not isinstance(video_dit_config, dict): | |
| raise ValueError(f"`video_dit_config` must resolve to a dict, got {type(video_dit_config)}") | |
| if isinstance(action_dit_config, DictConfig): | |
| action_dit_config = OmegaConf.to_container(action_dit_config, resolve=True) | |
| if action_dit_config is None: | |
| action_dit_config = {} | |
| if not isinstance(action_dit_config, dict): | |
| raise ValueError(f"`action_dit_config` must resolve to a dict, got {type(action_dit_config)}") | |
| if isinstance(video_scheduler, DictConfig): | |
| video_scheduler = OmegaConf.to_container(video_scheduler, resolve=True) | |
| if video_scheduler is None: | |
| video_scheduler = {} | |
| if not isinstance(video_scheduler, dict): | |
| raise ValueError(f"`video_scheduler` must be dict-like, got {type(video_scheduler)}") | |
| if isinstance(action_scheduler, DictConfig): | |
| action_scheduler = OmegaConf.to_container(action_scheduler, resolve=True) | |
| if action_scheduler is None: | |
| raise ValueError("`action_scheduler` is required for FastWAM.") | |
| if not isinstance(action_scheduler, dict): | |
| raise ValueError(f"`action_scheduler` must be dict-like, got {type(action_scheduler)}") | |
| required_action_scheduler_keys = {"train_shift", "infer_shift", "num_train_timesteps"} | |
| missing_keys = required_action_scheduler_keys - set(action_scheduler.keys()) | |
| if missing_keys: | |
| raise ValueError( | |
| f"`action_scheduler` missing required keys: {sorted(missing_keys)}. " | |
| "Expected keys: train_shift, infer_shift, num_train_timesteps." | |
| ) | |
| if isinstance(loss, DictConfig): | |
| loss = OmegaConf.to_container(loss, resolve=True) | |
| if loss is None: | |
| loss = {} | |
| if not isinstance(loss, dict): | |
| raise ValueError(f"`loss` must be dict-like, got {type(loss)}") | |
| return FastWAMIDM.from_wan22_pretrained( | |
| device=device, | |
| torch_dtype=model_dtype, | |
| model_id=model_id, | |
| tokenizer_model_id=tokenizer_model_id, | |
| video_cond_noise_prob=float(video_cond_noise_prob), | |
| tokenizer_max_len=int(tokenizer_max_len), | |
| load_text_encoder=bool(load_text_encoder), | |
| proprio_dim=(None if proprio_dim is None else int(proprio_dim)), | |
| redirect_common_files=bool(redirect_common_files), | |
| video_dit_config=video_dit_config, | |
| action_dit_config=action_dit_config, | |
| action_dit_pretrained_path=action_dit_pretrained_path, | |
| skip_dit_load_from_pretrain=bool(skip_dit_load_from_pretrain), | |
| mot_checkpoint_mixed_attn=bool(mot_checkpoint_mixed_attn), | |
| video_train_shift=float(video_scheduler.get("train_shift", 5.0)), | |
| video_infer_shift=float(video_scheduler.get("infer_shift", 5.0)), | |
| video_num_train_timesteps=int(video_scheduler.get("num_train_timesteps", 1000)), | |
| action_train_shift=float(action_scheduler["train_shift"]), | |
| action_infer_shift=float(action_scheduler["infer_shift"]), | |
| action_num_train_timesteps=int(action_scheduler["num_train_timesteps"]), | |
| loss_lambda_video=float(loss.get("lambda_video", 1.0)), | |
| loss_lambda_action=float(loss.get("lambda_action", 1.0)), | |
| compile_training_denoise=bool(compile_training_denoise), | |
| ) | |
| def create_fastwam_optional_idm( | |
| model_id: str, | |
| tokenizer_model_id: str, | |
| video_dit_config, | |
| action_idm_prob: float, | |
| video_cond_noise_prob: float = 0.5, | |
| tokenizer_max_len: int = 512, | |
| load_text_encoder: bool = True, | |
| proprio_dim: int | None = None, | |
| action_dit_config=None, | |
| action_dit_pretrained_path: str | None = None, | |
| skip_dit_load_from_pretrain: bool = False, | |
| video_scheduler=None, | |
| action_scheduler=None, | |
| loss=None, | |
| mot_checkpoint_mixed_attn: bool = False, | |
| compile_training_denoise: bool = False, | |
| redirect_common_files: bool = True, | |
| model_dtype: torch.dtype = torch.bfloat16, | |
| device: str = "cuda", | |
| ): | |
| from .models.wan22.fastwam_optional_idm import FastWAMOptionalIDM | |
| if isinstance(video_dit_config, DictConfig): | |
| video_dit_config = OmegaConf.to_container(video_dit_config, resolve=True) | |
| if not isinstance(video_dit_config, dict): | |
| raise ValueError(f"`video_dit_config` must resolve to a dict, got {type(video_dit_config)}") | |
| if isinstance(action_dit_config, DictConfig): | |
| action_dit_config = OmegaConf.to_container(action_dit_config, resolve=True) | |
| if action_dit_config is None: | |
| action_dit_config = {} | |
| if not isinstance(action_dit_config, dict): | |
| raise ValueError(f"`action_dit_config` must resolve to a dict, got {type(action_dit_config)}") | |
| if isinstance(video_scheduler, DictConfig): | |
| video_scheduler = OmegaConf.to_container(video_scheduler, resolve=True) | |
| if video_scheduler is None: | |
| video_scheduler = {} | |
| if not isinstance(video_scheduler, dict): | |
| raise ValueError(f"`video_scheduler` must be dict-like, got {type(video_scheduler)}") | |
| if isinstance(action_scheduler, DictConfig): | |
| action_scheduler = OmegaConf.to_container(action_scheduler, resolve=True) | |
| if action_scheduler is None: | |
| raise ValueError("`action_scheduler` is required for FastWAMOptionalIDM.") | |
| if not isinstance(action_scheduler, dict): | |
| raise ValueError(f"`action_scheduler` must be dict-like, got {type(action_scheduler)}") | |
| required_action_scheduler_keys = {"train_shift", "infer_shift", "num_train_timesteps"} | |
| missing_keys = required_action_scheduler_keys - set(action_scheduler.keys()) | |
| if missing_keys: | |
| raise ValueError( | |
| f"`action_scheduler` missing required keys: {sorted(missing_keys)}. " | |
| "Expected keys: train_shift, infer_shift, num_train_timesteps." | |
| ) | |
| if isinstance(loss, DictConfig): | |
| loss = OmegaConf.to_container(loss, resolve=True) | |
| if loss is None: | |
| loss = {} | |
| if not isinstance(loss, dict): | |
| raise ValueError(f"`loss` must be dict-like, got {type(loss)}") | |
| return FastWAMOptionalIDM.from_wan22_pretrained( | |
| device=device, | |
| torch_dtype=model_dtype, | |
| model_id=model_id, | |
| tokenizer_model_id=tokenizer_model_id, | |
| action_idm_prob=float(action_idm_prob), | |
| video_cond_noise_prob=float(video_cond_noise_prob), | |
| tokenizer_max_len=int(tokenizer_max_len), | |
| load_text_encoder=bool(load_text_encoder), | |
| proprio_dim=(None if proprio_dim is None else int(proprio_dim)), | |
| redirect_common_files=bool(redirect_common_files), | |
| video_dit_config=video_dit_config, | |
| action_dit_config=action_dit_config, | |
| action_dit_pretrained_path=action_dit_pretrained_path, | |
| skip_dit_load_from_pretrain=bool(skip_dit_load_from_pretrain), | |
| mot_checkpoint_mixed_attn=bool(mot_checkpoint_mixed_attn), | |
| video_train_shift=float(video_scheduler.get("train_shift", 5.0)), | |
| video_infer_shift=float(video_scheduler.get("infer_shift", 5.0)), | |
| video_num_train_timesteps=int(video_scheduler.get("num_train_timesteps", 1000)), | |
| action_train_shift=float(action_scheduler["train_shift"]), | |
| action_infer_shift=float(action_scheduler["infer_shift"]), | |
| action_num_train_timesteps=int(action_scheduler["num_train_timesteps"]), | |
| loss_lambda_video=float(loss.get("lambda_video", 1.0)), | |
| loss_lambda_action=float(loss.get("lambda_action", 1.0)), | |
| compile_training_denoise=bool(compile_training_denoise), | |
| ) | |
| def build_datasets(data_cfg: DictConfig): | |
| train_ds = instantiate(data_cfg.train) | |
| if data_cfg.get("val") is None: | |
| val_ds = train_ds | |
| else: | |
| train_stats_path = data_cfg.train.get("pretrained_norm_stats") | |
| default_stats_path = os.path.join(misc.get_work_dir(), "dataset_stats.json") | |
| val_stats_path = data_cfg.val.get("pretrained_norm_stats") | |
| pretrained_norm_stats = val_stats_path or train_stats_path or default_stats_path | |
| logger.info("Building val dataset with pretrained_norm_stats: %s", pretrained_norm_stats) | |
| val_ds = instantiate(data_cfg.val, pretrained_norm_stats=pretrained_norm_stats) | |
| return train_ds, val_ds | |
| def _resolve_train_device() -> str: | |
| if not torch.cuda.is_available(): | |
| return "cpu" | |
| device_count = torch.cuda.device_count() | |
| if device_count <= 1: | |
| return "cuda:0" | |
| local_rank = int(os.environ.get("LOCAL_RANK", "0")) | |
| if local_rank < 0 or local_rank >= device_count: | |
| return "cuda:0" | |
| return f"cuda:{local_rank}" | |
| def run_inference(cfg: DictConfig): | |
| setup_logging(log_level=logging.INFO) | |
| inference_cfg = cfg.inference | |
| mixed_precision = _normalize_mixed_precision(cfg.mixed_precision) | |
| model_dtype = _mixed_precision_to_model_dtype(mixed_precision) | |
| model = instantiate(cfg.model, model_dtype=model_dtype, device=str(inference_cfg.device)) | |
| checkpoint_path = inference_cfg.get("checkpoint_path") | |
| if checkpoint_path: | |
| ckpt = Path(checkpoint_path) | |
| if ckpt.exists(): | |
| logger.info("Loading finetuned checkpoint: %s", checkpoint_path) | |
| model.load_checkpoint(checkpoint_path) | |
| else: | |
| logger.warning("Checkpoint not found, skipping load: %s", checkpoint_path) | |
| model.eval() | |
| def center_crop_resize(img: Image, width: int, height: int) -> Image.Image: | |
| src_w, src_h = img.size | |
| scale = max(width / src_w, height / src_h) | |
| resized = img.resize((round(src_w * scale), round(src_h * scale)), resample=Image.BILINEAR) | |
| rw, rh = resized.size | |
| left = max((rw - width) // 2, 0) | |
| top = max((rh - height) // 2, 0) | |
| return resized.crop((left, top, left + width, top + height)) | |
| input_image = Image.open(str(inference_cfg.input_image_path)).convert("RGB") | |
| input_image = center_crop_resize(input_image, width=inference_cfg.width, height=inference_cfg.height) | |
| arr = np.array(input_image, dtype=np.float32) | |
| x = torch.from_numpy(arr) | |
| x = x.to(device=model.device, dtype=model.torch_dtype) | |
| x = x * (2.0 / 255.0) - 1.0 | |
| x = repeat(x, "H W C -> B C H W", B=1) | |
| output_mp4 = str(inference_cfg.output_mp4) | |
| infer_kwargs = { | |
| "prompt": str(inference_cfg.prompt), | |
| "negative_prompt": str(inference_cfg.negative_prompt), | |
| "text_cfg_scale": float(inference_cfg.text_cfg_scale), | |
| "action_cfg_scale": float(inference_cfg.action_cfg_scale), | |
| "input_image": x, | |
| "num_frames": int(inference_cfg.num_frames), | |
| "num_inference_steps": int(inference_cfg.num_inference_steps), | |
| "sigma_shift": None if inference_cfg.get("sigma_shift") is None else float(inference_cfg.sigma_shift), | |
| "seed": int(inference_cfg.seed), | |
| "rand_device": str(inference_cfg.rand_device), | |
| "tiled": bool(inference_cfg.tiled), | |
| } | |
| infer_out = model.infer(**infer_kwargs) | |
| video = infer_out["video"] | |
| save_mp4(video, output_mp4, fps=15) | |
| logger.info("Saved inference video to %s", output_mp4) | |
| return output_mp4 | |