refactor(seedvr2): split loader and sampler

This commit is contained in:
gaclove
2026-08-31 17:21:26 +08:00
parent fc1495f8a0
commit 6713097eb8
3 changed files with 312 additions and 303 deletions
+94
View File
@@ -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 🎥🅥🅗🅢"
}
}
}
+5 -3
View File
@@ -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",
+213 -300
View File
@@ -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": "<tensor>", # 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": "<tensor>",
"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,)