diff --git a/example_workflows/test_seedvr.json b/example_workflows/test_seedvr.json new file mode 100644 index 0000000..91cba5d --- /dev/null +++ b/example_workflows/test_seedvr.json @@ -0,0 +1,94 @@ +{ + "1": { + "inputs": { + "video": "1954627330480766977_wan2-2.mp4", + "force_rate": 0, + "custom_width": 0, + "custom_height": 0, + "frame_load_cap": 0, + "skip_first_frames": 0, + "select_every_nth": 1, + "format": "AnimateDiff" + }, + "class_type": "VHS_LoadVideo", + "_meta": { + "title": "Load Video (Upload) 🎥🅥🅗🅢" + } + }, + "4": { + "inputs": { + "ckpt_name": "seedvr2_ema_3b_fp8.safetensors", + "precision": "fp8-sgl", + "cpu_offload": false, + "use_tiling_vae": true + }, + "class_type": "LightX2VSeedVR2Loader", + "_meta": { + "title": "LightX2V SeedVR2 Loader" + } + }, + "5": { + "inputs": { + "target_height": 1920, + "target_width": 1080, + "infer_steps": 1, + "segment_length": 81, + "segment_overlap": 1, + "seed": 3816942657, + "source_fps": 16, + "model": [ + "4", + 0 + ], + "images": [ + "1", + 0 + ] + }, + "class_type": "LightX2VSeedVR2Sampler", + "_meta": { + "title": "LightX2V SeedVR2 Sampler" + } + }, + "6": { + "inputs": { + "video_info": [ + "1", + 3 + ] + }, + "class_type": "VHS_VideoInfo", + "_meta": { + "title": "Video Info 🎥🅥🅗🅢" + } + }, + "7": { + "inputs": { + "frame_rate": [ + "6", + 0 + ], + "loop_count": 0, + "filename_prefix": "AnimateDiff", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "trim_to_audio": false, + "pingpong": false, + "save_output": true, + "images": [ + "5", + 0 + ], + "audio": [ + "1", + 2 + ] + }, + "class_type": "VHS_VideoCombine", + "_meta": { + "title": "Video Combine 🎥🅥🅗🅢" + } + } + } \ No newline at end of file diff --git a/nodes/__init__.py b/nodes/__init__.py index 8b79c7b..32a1e75 100644 --- a/nodes/__init__.py +++ b/nodes/__init__.py @@ -21,7 +21,7 @@ from .config import ( ) from .inference import LightX2VModularInferenceV2 from .lora import LightX2VLoRALoader -from .seedvr import LightX2VSeedVRSR +from .seedvr import LightX2VSeedVR2Loader, LightX2VSeedVR2Sampler from .talk import ( TalkObjectInput, TalkObjectsCombiner, @@ -38,7 +38,8 @@ NODE_CLASS_MAPPINGS = { "LightX2VConfigCombinerV2": LightX2VConfigCombinerV2, "LightX2VConfigCombinerV3": LightX2VConfigCombinerV3, "LightX2VModularInferenceV2": LightX2VModularInferenceV2, - "LightX2VSeedVRSR": LightX2VSeedVRSR, + "LightX2VSeedVR2Loader": LightX2VSeedVR2Loader, + "LightX2VSeedVR2Sampler": LightX2VSeedVR2Sampler, "LightX2VTalkObjectInput": TalkObjectInput, "LightX2VTalkObjectsCombiner": TalkObjectsCombiner, "LightX2VTalkObjectsFromJSON": TalkObjectsFromJSON, @@ -54,7 +55,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "LightX2VConfigCombinerV2": "LightX2V Config Combiner V2", "LightX2VConfigCombinerV3": "LightX2V Config Combiner V3", "LightX2VModularInferenceV2": "LightX2V Modular Inference V2", - "LightX2VSeedVRSR": "LightX2V SeedVR2 Super-Resolution", + "LightX2VSeedVR2Loader": "LightX2V SeedVR2 Loader", + "LightX2VSeedVR2Sampler": "LightX2V SeedVR2 Sampler", "LightX2VTalkObjectInput": "LightX2V Talk Object Input (Single)", "LightX2VTalkObjectsCombiner": "LightX2V Talk Objects Combiner", "LightX2VTalkObjectsFromFiles": "LightX2V Talk Objects From Files", diff --git a/nodes/seedvr.py b/nodes/seedvr.py index 52b6593..cbafae2 100644 --- a/nodes/seedvr.py +++ b/nodes/seedvr.py @@ -1,339 +1,252 @@ -"""SeedVR2 super-resolution node.""" +"""SeedVR2 super-resolution nodes for ComfyUI. +Split into: +- LightX2VSeedVR2Loader: pick a SeedVR2 DiT checkpoint under + models/lightx2v/seedvr2/, load it into VRAM, return a SEEDVR_MODEL handle. +- LightX2VSeedVR2Sampler: takes SEEDVR_MODEL + IMAGE + per-call params, + returns upscaled IMAGE frames. + +The sampler installs a small shim on the runner so input frames come from the +IMAGE tensor (no temp file, no re-encode); the runner's segmenting logic still +runs and slices our in-memory tensor. +""" + +import argparse import gc -import hashlib import logging +import math +import types +from pathlib import Path +import folder_paths import torch from comfy.utils import ProgressBar -from ..model_utils import get_model_base_path, get_model_full_path, scan_models +logger = logging.getLogger(__name__) -class LightX2VSeedVRSR: - """SeedVR2 video/image super-resolution node for ComfyUI. +def _seedvr2_model_dir() -> Path: + return Path(folder_paths.models_dir) / "lightx2v" / "seedvr2" - Wraps the SeedVR2-3B model via LightX2V's SeedVRRunner to perform - single-pass diffusion super-resolution on a video (mp4) or a single image. + +def _scan_seedvr2_ckpts(): + d = _seedvr2_model_dir() + if not d.exists(): + return ["None"] + items = sorted(f.name for f in d.iterdir() if f.is_file() and f.suffix == ".safetensors") + return items or ["None"] + + +def _install_tensor_input_shim(runner, frames_u8, fps): + """Patch runner methods so input frames come from `frames_u8` instead of disk. + + frames_u8: torch.uint8 [T, C, H, W] on CPU (same format as torchvision.io.read_video output). """ + runner._tensor_input = frames_u8 + runner._tensor_input_fps = float(fps) - _current_runner = None - _current_config_hash = None + def _probe_video(self, video_path): # noqa: ARG001 + total = self._tensor_input.shape[0] + self._set_output_fps(self._tensor_input_fps) + return total, self._tensor_input_fps, None + + def _read_video_segment(self, video_path, start_idx, end_idx): # noqa: ARG001 + seg = self._tensor_input[start_idx:end_idx] + if seg.shape[0] == 0: + return torch.empty(0, 3, 0, 0, dtype=torch.uint8) + return seg + + original_encoder = runner._run_input_encoder_local_sr.__func__ + + def _run_input_encoder_local_sr(self): + if getattr(self, "_sr_segment", None) is None: + self._sr_segment = (0, self._tensor_input.shape[0]) + try: + return original_encoder(self) + finally: + self._sr_segment = None + return original_encoder(self) + + runner._probe_video = types.MethodType(_probe_video, runner) + runner._read_video_segment = types.MethodType(_read_video_segment, runner) + runner._run_input_encoder_local_sr = types.MethodType(_run_input_encoder_local_sr, runner) + runner.run_input_encoder = runner._run_input_encoder_local_sr + + +class LightX2VSeedVR2Loader: + """Load a SeedVR2 DiT checkpoint from models/lightx2v/seedvr2/.""" @classmethod def INPUT_TYPES(cls): - available_models = scan_models() + ckpts = _scan_seedvr2_ckpts() return { "required": { - "model_name": ( - available_models, + "ckpt_name": ( + ckpts, + {"default": ckpts[0], "tooltip": "DiT .safetensors under models/lightx2v/seedvr2/"}, + ), + "precision": ( + ["auto", "bf16", "fp8-sgl", "fp8-q8f", "fp8-vllm"], { - "default": available_models[0] if available_models else "None", - "tooltip": "SeedVR2 model directory under models/lightx2v/", - }, - ), - "input_type": ( - ["video", "image"], - {"default": "video", "tooltip": "Whether to SR a video file or a single image"}, - ), - "input_path": ( - "STRING", - { - "default": "", - "tooltip": "Absolute path to input .mp4 (for video) or .png/.jpg (for image). For video, also accepts a directory of frames.", - }, - ), - "sr_ratio": ( - "FLOAT", - { - "default": 2.0, - "min": 1.0, - "max": 8.0, - "step": 0.5, - "tooltip": "Super-resolution ratio (e.g. 2.0 = 2x, 4.0 = 4x)", - }, - ), - "target_height": ( - "INT", - { - "default": 720, - "min": 64, - "max": 4096, - "step": 8, - "tooltip": "Output frame height (SeedVR NaDiT processes at native resolution)", - }, - ), - "target_width": ( - "INT", - { - "default": 1280, - "min": 64, - "max": 4096, - "step": 8, - "tooltip": "Output frame width (must be divisible by 16 for VAE)", - }, - ), - "fps": ( - "FLOAT", - { - "default": 16.0, - "min": 1.0, - "max": 60.0, - "step": 0.5, - "tooltip": "Output FPS for video SR (input video FPS is preserved if available)", - }, - ), - "segment_length": ( - "INT", - { - "default": 81, - "min": 16, - "max": 256, - "step": 1, - "tooltip": "Frames per segment for long video SR (1-step diffusion per segment)", - }, - ), - "segment_overlap": ( - "INT", - { - "default": 1, - "min": 0, - "max": 32, - "step": 1, - "tooltip": "Overlap frames between segments to prevent seams", - }, - ), - "seed": ( - "INT", - { - "default": 42, - "min": -1, - "max": 2**32 - 1, - "tooltip": "Random seed, -1 for random", - }, - ), - "prompt": ( - "STRING", - { - "default": "", - "multiline": True, - "tooltip": "Optional text prompt to guide detail synthesis (SeedVR uses pre-computed embeddings; prompt mostly affects style)", - }, - ), - "negative_prompt": ( - "STRING", - { - "default": "", - "multiline": True, - "tooltip": "Negative prompt for guidance", - }, - ), - "save_output": ( - "BOOLEAN", - { - "default": False, - "tooltip": "If True, also write the SR result to disk in addition to returning IMAGE tensor", - }, - ), - "output_path": ( - "STRING", - { - "default": "", - "tooltip": "Where to save (only used if save_output=True). Leave empty for auto-generated name next to input.", - }, - ), - "unload_after_inference": ( - "BOOLEAN", - { - "default": False, - "tooltip": "Unload SeedVR runner from VRAM after inference (frees ~6GB+ for other nodes)", + "default": "auto", + "tooltip": "auto = bf16 for fp16/bf16 weights, fp8-sgl for fp8 weights. fp8-sgl needs sgl-kernel (H100/SM90); fp8-q8f is the 4090 path.", }, ), + "cpu_offload": ("BOOLEAN", {"default": False, "tooltip": "Offload DiT blocks to CPU between forwards (slower; only needed on small VRAM)"}), + "use_tiling_vae": ("BOOLEAN", {"default": True, "tooltip": "Tile VAE to reduce peak memory"}), + } + } + + RETURN_TYPES = ("SEEDVR_MODEL",) + RETURN_NAMES = ("model",) + FUNCTION = "load" + CATEGORY = "LightX2V/SeedVR" + + def load(self, ckpt_name, precision, cpu_offload, use_tiling_vae): + from ..lightx2v.lightx2v.infer import init_runner + from ..lightx2v.lightx2v.utils.set_config import set_config + + model_dir = _seedvr2_model_dir() + if ckpt_name == "None": + raise FileNotFoundError(f"No .safetensors checkpoints found in {model_dir}") + + for required in ("ema_vae.pth", "pos_emb.pt", "neg_emb.pt"): + p = model_dir / required + if not p.is_file(): + raise FileNotFoundError( + f"Missing {p}. SeedVR2 needs VAE + pre-computed text embeddings (pos_emb.pt / neg_emb.pt) in the same directory as the DiT checkpoint." + ) + ckpt_path = model_dir / ckpt_name + if not ckpt_path.is_file(): + raise FileNotFoundError(str(ckpt_path)) + + if precision == "auto": + precision = "fp8-sgl" if "fp8" in ckpt_name.lower() else "bf16" + + config = { + "model_cls": "seedvr2", + "task": "sr", + "model_path": str(model_dir), + "infer_steps": 1, + "fps": 16, + "target_video_length": 81, + "target_height": 1080, + "target_width": 1920, + "use_tiling_vae": bool(use_tiling_vae), + "cpu_offload": bool(cpu_offload), + } + if precision.startswith("fp8-"): + config["dit_quantized_ckpt"] = str(ckpt_path) + config["dit_quant_scheme"] = precision + config["dit_quantized"] = True + else: + config["dit_original_ckpt"] = str(ckpt_path) + + formatted = set_config(argparse.Namespace(**config)) + runner = init_runner(formatted) + logger.info(f"[SeedVR2Loader] loaded {ckpt_name} ({precision}); cpu_offload={cpu_offload}, tile_vae={use_tiling_vae}") + return ({"runner": runner, "precision": precision, "ckpt": ckpt_name},) + + +class LightX2VSeedVR2Sampler: + """Run SeedVR2 SR on an input frame tensor; return upscaled frames as IMAGE.""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("SEEDVR_MODEL",), + "images": ("IMAGE",), + "target_height": ("INT", {"default": 1080, "min": 64, "max": 4320, "step": 8, "tooltip": "Target output frame height. NaDiT preserves input aspect ratio; the geometric mean of target_h * target_w is the effective resolution cap."}), + "target_width": ("INT", {"default": 1920, "min": 64, "max": 7680, "step": 8}), + "infer_steps": ("INT", {"default": 1, "min": 1, "max": 50}), + "segment_length": ("INT", {"default": 81, "min": 16, "max": 512, "step": 1, "tooltip": "Frames per SR pass. Long videos are auto-segmented."}), + "segment_overlap": ("INT", {"default": 1, "min": 0, "max": 32}), + "seed": ("INT", {"default": 42, "min": 0, "max": 2**32 - 1}), + "source_fps": ("FLOAT", {"default": 16.0, "min": 1.0, "max": 120.0, "step": 0.5, "tooltip": "FPS of the input frames (passed through to the runner for any internal timing logic)"}), } } RETURN_TYPES = ("IMAGE",) RETURN_NAMES = ("images",) - FUNCTION = "run_seedvr" + FUNCTION = "sample" CATEGORY = "LightX2V/SeedVR" - def _config_hash( - self, - model_name, - input_type, - input_path, - sr_ratio, - target_height, - target_width, - fps, - segment_length, - segment_overlap, - seed, - prompt, - negative_prompt, - save_output, - output_path, - ): - """Hash of all parameters that should trigger runner reinit.""" - raw = ( - f"{model_name}|{input_type}|{input_path}|{sr_ratio}|" - f"{target_height}|{target_width}|{fps}|" - f"{segment_length}|{segment_overlap}|{seed}|" - f"{prompt}|{negative_prompt}|{save_output}|{output_path}" - ) - return hashlib.md5(raw.encode("utf-8")).hexdigest() + def sample(self, model, images, target_height, target_width, + infer_steps, segment_length, segment_overlap, seed, source_fps): + from ..lightx2v.lightx2v.utils.input_info import init_empty_input_info, update_input_info_from_dict - def run_seedvr( - self, - model_name, - input_type, - input_path, - sr_ratio, - target_height, - target_width, - fps, - segment_length, - segment_overlap, - seed, - prompt, - negative_prompt, - save_output, - output_path, - unload_after_inference, - ): - """Run SeedVR2 super-resolution and return IMAGE tensor.""" - from ..lightx2v.lightx2v.infer import init_runner - from ..lightx2v.lightx2v.utils.input_info import ( - init_empty_input_info, - update_input_info_from_dict, - ) - from ..lightx2v.lightx2v.utils.set_config import set_config + runner = model["runner"] - if not model_name or model_name == "None": - raise ValueError("model_name is required — select a SeedVR2 model directory under models/lightx2v/") - if not input_path: - raise ValueError("input_path is required — provide an absolute path to a video or image file") + if images.dim() != 4 or images.shape[-1] not in (3, 4): + raise ValueError(f"Expected IMAGE [T, H, W, C], got shape {tuple(images.shape)}") + # ComfyUI IMAGE: [T, H, W, C] float[0,1] → [T, C, H, W] uint8 (read_video's contract) + ori_h, ori_w = int(images.shape[1]), int(images.shape[2]) + frames = images[..., :3].permute(0, 3, 1, 2).contiguous() + frames_u8 = (frames.clamp(0.0, 1.0) * 255.0).to(torch.uint8).cpu() - model_full_path = get_model_full_path(model_name) - if not model_full_path: - raise FileNotFoundError( - f"Model '{model_name}' not found under models/lightx2v/. Expected directory: {get_model_base_path() / model_name}" + # Derive sr_ratio from input vs target. The runner uses + # resolution = min(sqrt(ori_h*ori_w) * sr_ratio, sqrt(target_h*target_w)) + # so we pick sr_ratio so the min lands on the target term (clamped to >=1 + # to avoid asking the SR model to downscale). + ori_geom = math.sqrt(ori_h * ori_w) + target_geom = math.sqrt(target_height * target_width) + sr_ratio = max(target_geom / ori_geom, 1.0) if ori_geom > 0 else 1.0 + if target_geom < ori_geom: + logger.warning( + f"[SeedVR2] target ({target_height}x{target_width}) smaller than input ({ori_h}x{ori_w}); SR will run at input scale." ) - cfg_hash = self._config_hash( - model_name, - input_type, - input_path, - sr_ratio, - target_height, - target_width, - fps, - segment_length, - segment_overlap, - seed, - prompt, - negative_prompt, - save_output, - output_path, + _install_tensor_input_shim(runner, frames_u8, source_fps) + + # runner.config is a LockableDict (locked after init); set_config uses temporarily_unlocked. + runner.set_config({ + "sr_ratio": float(sr_ratio), + "target_height": int(target_height), + "target_width": int(target_width), + "target_video_length": int(segment_length), # vestigial for SR; keep aligned with segment_length + "sr_segment_length": int(segment_length), + "sr_overlap": int(segment_overlap), + "infer_steps": int(infer_steps), + "seed": int(seed), + "fps": float(source_fps), + "video_path": "", # truthy sentinel so segmenting logic runs; shim bypasses file I/O + "image_path": "", + "prompt": "", + "negative_prompt": "", + "save_result_path": "", + "return_result_tensor": True, + }) + + input_info = init_empty_input_info("sr") + update_input_info_from_dict( + input_info, + { + "video_path": "", + "image_path": "", + "prompt": "", + "negative_prompt": "", + "seed": int(seed), + "sr_ratio": float(sr_ratio), + "save_result_path": "", + "return_result_tensor": True, + }, ) + progress = ProgressBar(100) + if hasattr(runner, "set_progress_callback"): + runner.set_progress_callback(lambda cur, _tot: progress.update_absolute(cur)) + try: - needs_reinit = ( - getattr(self.__class__, "_current_runner", None) is None or getattr(self.__class__, "_current_config_hash", None) != cfg_hash - ) - - if needs_reinit: - if getattr(self.__class__, "_current_runner", None) is not None: - del self.__class__._current_runner - torch.cuda.empty_cache() - gc.collect() - - config = { - "model_cls": "seedvr2", - "task": "sr", - "model_path": model_full_path, - "sr_ratio": float(sr_ratio), - "target_height": int(target_height), - "target_width": int(target_width), - "target_video_length": int(segment_length), - "sr_segment_length": int(segment_length), - "sr_overlap": int(segment_overlap), - "fps": float(fps), - "infer_steps": 1, - "seed": int(seed), - "prompt": prompt, - "negative_prompt": negative_prompt, - } - - formatted_config = set_config(config) - self.__class__._current_runner = init_runner(formatted_config) - self.__class__._current_config_hash = cfg_hash - - runner = self.__class__._current_runner - - progress = ProgressBar(100) - - def _update_progress(current_step, _total): - progress.update_absolute(current_step) - - if hasattr(runner, "set_progress_callback"): - runner.set_progress_callback(_update_progress) - - input_info = init_empty_input_info("sr") - update_input_info_from_dict( - input_info, - { - "video_path": input_path if input_type == "video" else "", - "image_path": input_path if input_type == "image" else "", - "prompt": prompt, - "negative_prompt": negative_prompt, - "seed": int(seed), - "save_result_path": output_path if (save_output and output_path) else "", - "return_result_tensor": True, - }, - ) - - runner.set_config( - { - "video_path": input_path if input_type == "video" else "", - "image_path": input_path if input_type == "image" else "", - "prompt": prompt, - "negative_prompt": negative_prompt, - "seed": int(seed), - "save_result_path": output_path if (save_output and output_path) else "", - "return_result_tensor": True, - } - ) - - result_dict = runner.run_pipeline(input_info) - images = result_dict.get("video", None) - - if images is None or images.numel() == 0: - raise RuntimeError("SeedVR returned empty result") - - images = images.cpu() - if images.dtype != torch.float32: - images = images.float() - - if images.dim() == 4 and images.shape[0] > 0: - images = images[0] - - if unload_after_inference: - if hasattr(self.__class__, "_current_runner"): - del self.__class__._current_runner - self.__class__._current_runner = None - self.__class__._current_config_hash = None - + result = runner.run_pipeline(input_info) + finally: torch.cuda.empty_cache() gc.collect() - return (images,) + video = result.get("video") if isinstance(result, dict) else result + if video is None or video.numel() == 0: + raise RuntimeError("SeedVR2 returned empty result") - except Exception as e: - logging.error(f"SeedVR SR failed: {e}") - if unload_after_inference: - if hasattr(self.__class__, "_current_runner"): - del self.__class__._current_runner - self.__class__._current_runner = None - self.__class__._current_config_hash = None - raise + # wan_vae_to_comfy already gives [T, H, W, C] float[0,1] on CPU + video = video.detach().cpu().float().clamp(0.0, 1.0) + return (video,)