Skip to content

Releases: ussoewwin/ComfyUI-SeedVR2-VideoUpscaler-with-TensorRT

v1.6.5 — TensorRT VAE Fallback Elimination & Engine Build ONNX Auto-Cleanup

Choose a tag to compare

@ussoewwin ussoewwin released this 10 Oct 13:43
EN 中文

Overview

SeedVR2 Video Upscaler v1.6.5 resolves a critical stability failure in the VAE inference pipeline where unintended fallback to standard PyTorch FP16 VAE occurred during multi-stage or single-VAE workflow configurations, leading to catastrophic CUDA Out Of Memory (OOM) errors at high resolutions. Additionally, this release introduces automated post-build cleanup for intermediate ONNX graph artifacts in TensorRT engine builders.


1. Complete Elimination of Unintended PyTorch FP16 VAE Fallbacks

Root Cause & Vulnerability Analysis

In previous releases, SeedVR2VideoUpscaler exposed separate inputs for vae_encode and vae_decode. When workflows only wired one of the inputs (e.g. connecting a TensorRT VAE Loader solely to vae_decode while leaving vae_encode unconnected), encode_cfg evaluated to an empty dictionary ({}). Consequently:

  • runner.use_tensorrt_vae_encode evaluated to False.
  • A legacy flag decoupling policy (# TRT decoder + FP16 encoder must keep the encoder on the FP16 path) allowed the encoder to silently decouple from TensorRT.
  • Phase 1 (VAE Encoding) silently fell back to PyTorch FP16 VAE (VideoAutoencoderKLWrapper), completely bypassing the dedicated TensorRT 1-shot engine.

At high resolutions (such as 2160x2160px padded to 4208x4208px in 2-stage upscaling pipelines), PyTorch VAE's un-tiled 3D causal convolutions (InflatedCausalConv3d) attempted to allocate upwards of 30.69 GiB (26.47 GiB allocated + 4.22 GiB requested) during slice concatenation, immediately causing CUDA OOM on standard 16GB GPUs.

Technical Remediation

  1. Bidirectional Configuration Mirroring:
    If either vae_encode or vae_decode is provided, the node now automatically mirrors the configuration to both endpoints:
    encode_cfg = dict(vae_encode) if vae_encode is not None else (dict(vae_decode) if vae_decode is not None else {})
    decode_cfg = dict(vae_decode) if vae_decode is not None else (dict(vae_encode) if vae_encode is not None else {})
  2. Strict TensorRT Synchronization:
    If TensorRT is active or requested on either endpoint, both encode_cfg and decode_cfg are strictly locked to TensorRT:
    _is_trt = bool(
        encode_cfg.get("use_tensorrt_vae", False)
        or decode_cfg.get("use_tensorrt_vae", False)
        or encode_cfg.get("vae_backend") == "tensorrt"
        or decode_cfg.get("vae_backend") == "tensorrt"
    )
    if _is_trt:
        encode_cfg["use_tensorrt_vae"] = True
        encode_cfg["vae_backend"] = "tensorrt"
        decode_cfg["use_tensorrt_vae"] = True
        decode_cfg["vae_backend"] = "tensorrt"
    
    runner.use_tensorrt_vae_encode = _is_trt
    runner.use_tensorrt_vae_decode = _is_trt
    runner.use_tensorrt_vae = _is_trt
  3. Hardened Infer Checks (src/core/infer.py):
    In VideoDiffusionInfer.vae_encode() and vae_decode(), the TensorRT active checks (_enc_trt and _dec_trt) now verify all runner-level TRT indicators, guaranteeing that any active TRT workflow directly invokes the dedicated TensorRT engine with zero silent dropbacks to PyTorch FP16.

2. Automated Cleanup of Intermediate ONNX Artifacts upon Engine Build

During dedicated TensorRT RTX engine compilation via SeedVR2BuildTensorRTVAE (src/interfaces/trt_vae_builder.py) or dynamic loader builds (src/interfaces/trt_vae_model_loader.py):

  • Tracing large video batch graphs previously created temporary .onnx and companion .onnx.data files in tensorrt_backend/artifacts/.
  • While necessary as input for cloud_build_engine.py, these ONNX files are completely redundant once the final .rtxplan engine is built and verified.
  • Automated Immediate Purge: Immediately upon successful engine serialization, onnx_path and onnx_data are automatically unlinked and safely removed from disk. This prevents multi-gigabyte temporary graphs from cluttering disk storage.

3. Node Registration & Display Name Standardization

  • Cleaned up node registration display names in SeedVR2VideoUpscaler by removing legacy hardcoded version identifiers, standardizing to SeedVR2 Video Upscaler.
  • Preserved full backward compatibility with existing workflows and saved graph JSON schemas.

4. Modified Files & Verification

  • src/interfaces/video_upscaler.py: Bidirectional VAE config mirroring, strict TRT flag locking.
  • src/core/infer.py: Strict multi-flag TRT encoder/decoder checks.
  • src/interfaces/trt_vae_builder.py: Post-build ONNX and ONNX data unlinking.
  • src/interfaces/trt_vae_model_loader.py: Post-build ONNX unlinking in runtime build paths.
  • md/changelog.md & zhmd/changelogzh.md: Bilingual changelog synchronization.

v1.6.3

Choose a tag to compare

@ussoewwin ussoewwin released this 09 Oct 04:09
EN 中文

Overview

This release removes the hardcoded version pinning (triton-windows==3.5.1.post24) from requirements-windows-cu132.txt, transitioning Triton dependency resolution to dynamic latest package resolution.

By unpinning triton-windows, pip and the installer pipeline automatically resolve and install the latest available Windows wheel from PyPI, ensuring seamless compatibility with ongoing ComfyUI, PyTorch, and CUDA environment updates without manual version locks.


Changes

1. Requirements (requirements-windows-cu132.txt)

  • Removed hardcoded version pin triton-windows==3.5.1.post24 in favor of unpinned triton-windows.
  • Aligns requirements-windows-cu132.txt with install.py's dynamic resolution logic (ensure_package("triton", "triton-windows" if sys.platform.startswith("win") else "triton", no_deps=True)).

2. Documentation & Multilingual Sync

v1.6.2 — norm bf16 Mode on the Legacy DiT Loader

Choose a tag to compare

@ussoewwin ussoewwin released this 08 Oct 08:25
EN 中文

This release is the norm_bf16 mode landing on the standard (legacy) DiT loader. The norm_bf16
switch, previously exclusive to the DisTorch2 loader (v1.6.0), is now implemented on the standard
SeedVR2 (Down)Load DiT Model node as well, so both loader nodes expose identical
norm-precision control: the VRAM-saving bf16 norm path no longer requires the DisTorch2 loader.


1. norm_bf16 switch on the legacy DiT loader

Commit: 12c2534 — feat(loader): expose norm_bf16 switch on legacy DiT loader node

  • What it does: during DiT upscaling (Phase 2) the RMS/QK norm operations run on the stock fp32
    path by default, which keeps a notable amount of activation memory resident on the GPU.
    norm_bf16 = ON runs these norm operations in bf16 instead — saving significant resident
    VRAM during Phase 2, at the cost of bf16 rounding (measured per-pixel PSNR ~37–39 dB vs fp32).
  • OFF (default): stock fp32 norm path — quality-priority; behaviour unchanged from previous
    releases (the fp32 path code was not touched).
  • The DiT-side bf16 norm path itself is unchanged from v1.6.0; this release only opens the same
    path to the legacy loader node. No model-side code changes.

2. Dual-entry switch wiring with explicit separation

Commit: 12c2534 (same commit — the video_upscaler.py hunk)

  • The per-node switch in SeedVR2 Video Upscaler now reads both loader config layouts with
    explicit branch separation: DisTorch2 loader → distorch2.norm_bf16; legacy loader →
    top-level norm_bf16. The discriminator is the presence of the distorch2 block; the two
    entry points are never mixed-read.
  • Verified by executing the shipped switch block itself across 6 scenarios: legacy ON/OFF,
    DisTorch2 ON/OFF, stray top-level key with a DisTorch2 block present (d2 value wins), and
    missing key → OFF.
  • Verified the node schema contains the new norm_bf16 input and that execute() propagates it
    into the config with default OFF.

3. Node documentation

Commit: ebdefe4 — docs(readme): add SeedVR2 (Down)Load DiT Model (legacy loader) node section
with screenshot (EN/zh)

  • Both READMEs (EN + 中文) gained a full parameter section for the legacy loader node — it had no
    README section before — including the norm_bf16 toggle, with workflow screenshot
    docs/legacy_dit.png.

4. Summary of the norm_bf16 change set (v1.6.1 → v1.6.2)

Commit Change Effect
12c2534 norm_bf16 on the legacy loader + separated dual-entry switch bf16 norm path now on both loaders
ebdefe4 README node section (EN/zh) + screenshot docs
3d0a953 changelog v1.6.2 entry (EN/zh) docs

Net result: whichever loader node the workflow uses, norm_bf16 = ON gives the same Phase 2
VRAM-saving bf16 norm path; the default OFF leaves the stock fp32 path untouched.


5. Recommended use

  • Low-VRAM GPUs (e.g. 16 GB): ON is the recommended way to shave resident VRAM during Phase 2
    when close to the limit — pair it with the existing controls (blocks_to_swap,
    emb_repeat_nocache, or the DisTorch2 loader).
  • Quality-critical runs: keep it OFF (default) to stay on the fp32 norm path.

v1.6.1 — SpargeAttn Speed Work (Sage2++ / Batched Windows / Zero Sync / Varlen)

Choose a tag to compare

@ussoewwin ussoewwin released this 05 Oct 04:24
EN 中文

This release is the SpargeAttn speed work since v1.6.0. The spargeattn attention backend
(the block-sparse attention built on the SageAttention2++ kernels) was previously slower than
plain SageAttention2
on the live workflow, even though it is the evolved form. The four changes
below remove the fixed overheads that caused that, so spargeattn now meets or beats
SageAttention2 while keeping its sparsity gain.


1. Sage2++ fp16-accumulate path + ragged-window fix

Commit: 0185398 — perf: fix ragged-window SpargeAttn path and use Sage2++ fp16 accumulate

  • Ragged-window correctness: the per-group (n,H,L,D) tensors were built with view(), but
    q/k/v are laid out (total_seq,H,D) and the NA windowing produces ragged windows, so
    view(_w,L,H,D) raised "shape is invalid for input of size …". Now built with torch.stack,
    and outputs are written back per window.
  • Sage2++ fp16-accumulate kernel: the Stage-3 sparse kernel is driven through
    qk_int8_sv_f8_accum_f16_… (SageAttention2++ fp16-accumulate) when SAGE2PP_ENABLED is on and
    the symbol is present, falling back to f32-accumulate otherwise.

Measured (RTX 5060 Ti, sm120, bf16, H=24, D=128, 25×4032):

topk spargeattn (Sage2++) sageattn_2 Speedup
0.5 51.6 ms 70.3 ms 1.36x
0.25 42.6 ms 71.0 ms 1.67x

2. Equal-length window batching (stock API, zero-copy)

Commit: 8128ee6 — perf: batch SpargeAttn per equal-length window group (stock API, zero-copy)

The stock loop calls spas_sage2_attn_meansim_topk_cuda once per NA window. On the live
workflow (1080p, 21 frames → 100 windows/call, 36 blocks × 2 attn) that is ~7200 kernel
launches per step
vs sageattn_varlen's 72; WDDM launch overhead (~0.15 ms) eats the
sparsification gain and makes spargeattn slower than sageattn_2 even at low topk.

The fix groups windows by equal length and calls the stock topk API once per group on a
batched (n,H,L,D) tensor:

  • equal-length windows laid out back-to-back (the NA norm) use a zero-copy view;
  • ragged groups fall back to torch.cat.

The API and its internal block-map / LUT / quant / sparse kernels are natively batch-capable,
so this is the exact stock code path — no re-implementation. Outputs match the per-window
path (ragged 99×2430 + 1×1000: max_diff 1.5e-3, mean 4.4e-10).

Measured (RTX 5060 Ti, sm120, bf16, 100×2430):

topk spargeattn (batched) sageattn_2 Speedup
0.5 101.0 ms 111.9 ms 1.11x
0.3 90.6 ms 114.5 ms 1.26x

A batched block-map re-implementation was tried first and reverted (879441b): since
spas_sage2_attn_meansim_topk_cuda already runs one window per call and internally uses the
Sage2++ fp16-accumulate kernel, the re-implementation only added torch.stack copies and made
spargeattn slower. The batching above keeps the stock API and only changes the call granularity.


3. Per-call D2H sync elimination

Commit: 29a1d84 — perf: eliminate per-call D2H syncs in SpargeAttn fast path

The zero-copy uniform fast path ran a (seq == seq[0]).all() plus two .item() calls per
attention call
— that is 3 D2H syncs × 72 calls/DiT-forward × 14 batches. On Windows/WDDM
each sync busy-waits ~3–7 ms, adding seconds per batch and making spargeattn lose to plain
SageAttention2 (31s vs 29s).

The fix caches one plan per (queue-shape, total-numel) key: the first call pays the syncs
once, and every subsequent call in the run executes the zero-copy fast path with no D2H sync.
Stale plans cannot mis-execute — the zero-copy .view() validates B*L against the tensor numel
and pops the plan on mismatch, falling through to the ragged path.


4. Variable-length (varlen) delegation to the SpargeAttn-hswq library

Commit: 3d93819 (ef32a90 restores the measured baseline) — delegate varlen handling to the library

call_sparge_attn_varlen now delegates to the fork's spas_sage2_attn_meansim_topk_varlen_cuda
(≥ 1.2.1)
instead of doing caller-side bucketing. The library supports both layouts without ever
violating the kernels' fixed-length premise:

  • uniform (all windows the same length) → a single zero-copy batched launch;
  • mixed lengths → bucketed by identical (L_q, L_k) inside the library; each bucket costs
    one batched launch. Contiguous buckets use zero-copy views, scattered buckets use one gather +
    one scatter. No per-sequence launches, no padding waste.
  • the bucket plan is cached by the order-independent multiset of window lengths, so runs that
    permute the same window lengths between calls still hit it.

Note on the persistent-cache experiment: an experimental persistent Triton-cache feature
(TRITON_CACHE_DIR default + release_sparge_kernel_caches) was added and then reverted —
measured, it did not change steady-state speed and only made the first call slower. The spargeattn
path is left in its measured-correct state (no persistent-cache instrumentation).


5. Summary of the change set (v1.6.0 → v1.6.1)

Commit Change Measured effect
0185398 Sage2++ fp16-accumulate + ragged fix up to 1.67x vs SA2
879441b revert batched block-map re-implementation removes stack copies
8128ee6 equal-length window batching (stock API, zero-copy) 1.11–1.26x vs SA2
29a1d84 per-call D2H sync elimination removes 3 syncs × 72 calls/batch
3d93819 varlen delegation to library ≥ 1.2.1 uniform 1 launch / mixed bucketed
ef32a90 restore the measured-correct spargeattn baseline —

Net result: on the live workflow the spargeattn backend is now at parity with (or faster
than) SageAttention2, with the sparsity gain intact and no per-call host synchronization in the
fast path.


6. Library dependency

The varlen path requires SpargeAttn-hswq ≥ v1.2.1
(spas_sage2_attn_meansim_topk_varlen_cuda). See the
SpargeAttn-hswq v1.2.1 release.

v1.6.0 — DisTorch2 DiT Loader Node & Phase 2 VRAM Controls

Choose a tag to compare

@ussoewwin ussoewwin released this 03 Oct 23:36
EN 中文

This release adds a DisTorch2-backed DiT loader node that can host an entire quantized
(INT8 / NVFP4) DiT in system RAM and stream it to the compute device during denoising, plus two
new Phase 2 VRAM controls (norm_bf16, emb_repeat_nocache) and a hardening pass on the INT8
quantization path. It also trims the model registry (fp8 entries removed).

License note: the vendored DisTorch2 backend under src/distorch2/ is a verbatim copy of
ComfyUI-MultiGPU (pollockjj, GPL-3.0),
while the rest of this repository is Apache 2.0. See the README Credits/License sections.


1. New node — SeedVR2 (Down)Load DiT Model with Distorch2

  • node_id: SeedVR2LoadDiTModelDisTorch2
  • display_name: SeedVR2 (Down)Load DiT Model with Distorch2
  • Output: the same SEEDVR2_DIT config dict as the standard loader, plus a distorch2 block,
    so it connects to SeedVR2 Video Upscaler identically.

What it does

Hosts the whole packed DiT in a donor device (typically cpu = system RAM) and streams weights
to the compute device per block during the DiT forward, using the upstream DisTorch2 placement
backend. This lets a large quantized DiT run on a GPU with far less free VRAM than the model size.

Inputs

Input Type Default Purpose
model DiT checkpoint seedvr2_7b_int8_convrot.safetensors DiT model (INT8 / NVFP4 / FP16).
device device first device Compute device for DiT inference.
attention_mode combo sdpa sdpa / flash_attn_2 / flash_attn_3 / sageattn_2 / sageattn_3 / spargeattn.
sparge_topk string 0.5 SpargeAttn KV block keep ratio.
distorch2_enabled bool True Enable DisTorch2 placement; OFF behaves like the standard loader.
virtual_vram_gb float 4.0 Virtual-VRAM budget on the compute device for streaming (0 = keep whole model on donor).
donor_device combo cpu Device that physically hosts the packed DiT weights.
expert_mode_allocations string "" Advanced per-block device allocation string, e.g. "cpu,cpu,cuda:0".
eject_models bool True Eject other resident models before placement to free VRAM.
emb_repeat_nocache bool False Disable the Phase 2 emb_repeat cache (see §3).
norm_bf16 bool False bf16 norm path for Phase 2 (see §2).

The node logs its resolved placement in the console
([MultiGPU DisTorch V2] ... Final Allocation String and the per-device layer distribution table).

Allocation-string bridge (src/core/distorch2_placement.py)

The node produces an upstream-format allocation string
"<expert_mode_allocations>#<compute_device>;<virtual_vram_gb>;<donor_device>". The bridge
(apply_distorch2_placement) feeds the SeedVR2 DiT into the copied backend:

  • FakePatcherAdapter — presents the SeedVR2 DiT to analyze_safetensor_loading.
  • build_packed_allocation_string / _parse_vram_string — recompute the allocation from the
    packed weight sizes (the upstream logical-size calc otherwise assigns 0 to CPU).
  • wrap_blocks_with_cache_release — per-block cache release (the BlockSwap-era mechanism)
    when BlockSwap is not active.
  • ensure_backend_registered — idempotently registers the patched safetensor ModelPatcher.

Placement is applied before BlockSwap / compile wiring in materialize_model.


2. norm_bf16 — Phase 2 RMS/QK norm precision toggle

The DiT RMS/QK norm (CustomRMSNorm) is called many times per block. Under torch.autocast(bf16)
the stock path promotes intermediates to fp32, allocating full-size fp32 activations and
inflating Phase 2 VRAM.

  • OFF (default): stock fp32 path — output is bit-identical to the previous build.
  • ON: bf16 path — torch.mean(input * input, dim=..., dtype=input.dtype) and the weight cast
    to the input dtype, keeping the whole RMS in bf16.

The toggle is wired node → config dict → env var SEEDVR2_NORM_BF16, consumed inside the forward
(as with the existing switches). Applied to both 7B and 3B CustomRMSNorm.

Measured quality (this repo's test harness, RTX 5060 Ti): per-pixel PSNR ~37–39 dB vs the fp32
path (bf16 rounding), i.e. a deliberate VRAM-vs-quality trade.


3. emb_repeat_nocache — Phase 2 emb_repeat cache toggle

  • OFF (default): keep the emb_repeat cache (stock behaviour).
  • ON: recompute the repeat_interleave on every use. repeat_interleave is a pure copy, so the
    output is bit-identical; it frees the shared resident cache during Phase 2.

Wired node → config dict → env var SEEDVR2_EMB_REPEAT_NOCACHE.


4. INT8 quantization-path hardening (no implicit fp16)

Quantized DiT weights must stay packed (INT8 / uint8 storage) through the forward; several comfy
paths could silently fp16-expand them when activations were bf16. This release closes those paths
for the DiT only (VAE untouched):

  • src/core/model_loader.py: the DiT comfy_quant ops dtype is taken from the runner's pipeline
    compute dtype (bf16) instead of a hardcoded torch.float16, so non-quantized DiT params are not
    forced to fp16 and QuantizedTensor logical dtype is not tagged fp16.
  • src/optimization/int8_native_ops.py / nvfp4_native_ops.py: _mark_dit_locked wraps the
    DiT ops.Linear so packed (int8/uint8) weights are flagged _dit_quant_locked, which stops
    comfy_kitchen from silently fp16-expanding a locked weight on any unimplemented op.
  • comfy cast/dequant paths (cast_bias_weight, post_cast, forward_comfy_cast_weights,
    _load_quantized_module) are guarded so locked DiT weights are never dequantized to fp16.

5. Model registry & CLI

  • fp8 DiT entries removed from the auto-download registry (3B, 7B, 7B sharp). The remaining
    DiT entries are FP16, INT8 (ConvRot) and NVFP4.
  • DEFAULT_DIT is now seedvr2_7b_int8_convrot.safetensors.
  • CLI: new --distorch2 flags (--distorch2, --distorch2_virtual_vram_gb, --distorch2_donor,
    --distorch2_allocation) for headless use.

6. Documentation

  • README.md / zhmd/README.md: new "SeedVR2 (Down)Load DiT Model with Distorch2" node section
    with docs/distorch2.png and full parameter documentation.
  • Credits/License: DisTorch2 attribution in Credits, and the GPL-3.0 notice for the vendored
    backend in License (Apache 2.0 for the rest of this repository).
  • Changelog (EN/ZH) updated with this v1.6.0 entry.

Compatibility

  • The standard SeedVR2 (Down)Load DiT Model node is unchanged; the DisTorch2 node is additive.
  • norm_bf16 and emb_repeat_nocache default OFF, so default output is unchanged.
  • No change to the VAE or the TRT VAE paths in this release.

Full Changelog: v1.5.9...v1.6.0

v1.5.9 - SpargeAttn-hswq: New Block-Sparse Attention Backend for the DiT

Choose a tag to compare

@ussoewwin ussoewwin released this 02 Oct 04:31
EN 中文

SpargeAttn-hswq — New Block-Sparse Attention Backend for the DiT


1. Overview

v1.5.9 adds a new attention backend, spargeattn, built on the SpargeAttn-hswq fork (spas_sage_hswq_attn v1.0.0). It runs two-stage block-sparse filtering on top of SageAttention2++ quantized kernels (QK INT8 + PV FP8), using the fork's plug-and-play entry point spas_sage2_attn_meansim_topk_cuda — the "SpargeAttn + Sage2" combination from the paper's kernel benchmarks.

Attention is the dominant compute in the DiT: every transformer block attends over the packed vid+txt token sequence. SpargeAttn skips attention blocks whose softmax contribution is negligible, trading a user-selectable amount of accuracy for speed.

Attention Mode Kernel Precision Sparsity
sdpa PyTorch SDPA FP16/BF16 exact none
sageattn_2 SageAttention2 QK INT8 + PV FP8 none
spargeattn (new) SpargeAttn-hswq QK INT8 + PV FP8 + block skipping topK-selectable

Prebuilt Windows wheels: https://huggingface.co/ussoewwin/Sage-Attention-and-Sparge-Attention-HSWQ


2. Why Two-Stage Block-Sparse Filtering

SageAttention2++ already quantizes Q/K to INT8 and P/V to FP8 per block. SpargeAttn adds a block selection stage on top:

  1. Stage 1 (coarse): a low-precision attention pass (FP8 PV) computes approximate block importance scores across the KV grid.
  2. Stage 2 (topK): only the top-K fraction of KV blocks — measured cumulatively by their softmax contribution — is recomputed with the full-precision path. The remaining blocks are skipped entirely.

The keep ratio is exposed to the user as sparge_topk (1.0 = compute everything, smaller values skip more blocks and run faster).


3. Varlen-Aware Per-Window Execution

Target File: src/optimization/compatibility.py
Target Function: call_sparge_attn_varlen() (L606–684)

SeedVR2 does not call attention on one flat sequence. The vid and txt token streams are packed into a single varlen batch, and MSA / Swin-window attention both iterate over cu_seqlens windows. The SpargeAttn plug-and-play API has no native varlen support, so a per-window dispatch was implemented mirroring pytorch_varlen_attention():

def call_sparge_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k,
                            max_seqlen_q, max_seqlen_k,
                            sparge_topk: float = 0.5, **kwargs):
    # ...
    output = torch.empty_like(q)
    for i in range(num_seqs):
        q_start, q_end = int(cu_q_cpu[i]), int(cu_q_cpu[i + 1])
        k_start, k_end = int(cu_k_cpu[i]), int(cu_k_cpu[i + 1])
        q_i = q[q_start:q_end].permute(1, 0, 2).unsqueeze(0)  # (1, H, S, D)
        k_i = k[k_start:k_end].permute(1, 0, 2).unsqueeze(0)
        v_i = v[k_start:k_end].permute(1, 0, 2).unsqueeze(0)

        seq_len = q_i.size(-2)
        headdim = q_i.size(-1)
        if seq_len < 128 or headdim not in (64, 128):
            out_i = _sdpa(q_i, k_i, v_i, is_causal=is_causal)       # exact fallback
        else:
            out_i = spas_sage2_attn_meansim_topk_cuda(               # SpargeAttn kernel
                q_i, k_i, v_i, topk=_topk, is_causal=is_causal, tensor_layout="HND",
            )
        output[q_start:q_end] = out_i.squeeze(0).permute(1, 0, 2)

    return output.to(out_dtype) if output.dtype != out_dtype else output

Key implementation details:

Aspect Handling
Output buffer Pre-allocated torch.empty_like(q); each window writes its slice directly (no list accumulation, no final torch.cat)
Input packing (total_seq, heads, head_dim) → per-window (1, H, S, D) in "HND" tensor layout expected by the kernel
dtype Kernel requires FP16/BF16; inputs are normalized (K/V cast to Q's dtype; everything to BF16 if Q is FP32), output cast back to the caller's dtype at the end
Causality is_causal forwarded to both the kernel and the SDPA fallback

4. Per-Window Constraint Fallbacks (Never Silent)

The SpargeAttn kernel asserts hard constraints that SeedVR2's variable windows do not always satisfy. Each window is checked before kernel dispatch, and failing windows transparently fall back to exact SDPA:

Constraint Reason Fallback
seq_len >= 128 Kernel assert: block-sparse pooling requires at least 128 tokens per window Exact SDPA for that window
headdim ∈ {64, 128} Kernel assert (SeedVR2 uses 128; satisfied in practice) Exact SDPA for that window
topk ∉ (0, 1] Invalid keep ratio Sanitized to 0.5

Short windows occur at small input resolutions (spatial windows shrink below the 128-token floor). The fallback guarantees exactness wherever the kernel cannot run and sparsity wherever it can — the output is a per-window mix of exact and block-sparse attention, chosen deterministically by window shape.


5. Availability & Fallback Chain

Target Function: validate_attention_mode() (L273–298)

requested_mode = 'spargeattn'
    ├─ spas_sage_hswq_attn installed  → 'spargeattn'  (block-sparse kernels)
    ├─ else SageAttention2 installed  → 'sageattn_2'  (quantized, no sparsity)
    └─ else                           → 'sdpa'        (exact)

The fallback emits a logged WARNING with the wheel installation link — never silent. The same chain is enforced in FlashAttentionVarlen (the shared entry point), so CLI and node paths behave identically.


6. Configurable topK

Node Path (DiT loader)

New optional COMBO input sparge_topk on SeedVR2 Load DiT Model (src/interfaces/dit_model_loader.py L121–132):

  • Options: 0.05 … 1.0 in 0.05 steps, default 0.5
  • Higher = more KV blocks computed = more accurate, less acceleration; 1.0 = no skipping
  • Ignored by all other attention backends

CLI Path

--attention_mode spargeattn + --sparge_topk (inference_cli.py L1480–1483, same default and range).

Plumbing & Sanitization

src/core/model_configuration.py routes the value end-to-end and sanitizes it with _safe_topk(): anything outside (0, 1] (including non-numeric) falls back to 0.5. The sanitized value is applied to every attention module at model (re)configuration time and re-applied on cache hits, so switching models or attention modes never carries a stale topK.


7. Verification

GPU-tested on RTX 5060 Ti (SM120, torch 2.14.1+cu132):

  • Mixed-window varlen call (256 + 64 tokens): the 256-token window returns quantized SpargeAttn output; the sub-128 window falls back to SDPA and matches the exact reference with cosine similarity 1.000000.
  • End-to-end: FlashAttentionVarlen(attention_mode='spargeattn') exercised through the real DiT attention module (src/models/dit_3b/attention.py / dit_7b/attention.py), confirming the sparge_topk value reaches the kernel.
  • dit_3b and dit_7b attention modules remain byte-identical to each other in their SpargeAttn dispatch.

8. Files Changed

File Change
src/optimization/compatibility.py call_sparge_attn_varlen() (varlen per-window dispatch + fallbacks), SPARGE_ATTN_AVAILABLE detection, validate_attention_mode() spargeattn branch
src/models/dit_3b/attention.py spargeattn dispatch branch + sparge_topk parameter
src/models/dit_7b/attention.py Same (7B architecture)
src/interfaces/dit_model_loader.py sparge_topk COMBO input; attention_mode gains spargeattn option
src/core/model_configuration.py sparge_topk plumbing, _safe_topk() sanitization, cache-safe reapplication
src/core/generation_utils.py prepare_runner() signature passthrough
src/optimization/memory_manager.py Attention-mode description entry
inference_cli.py --attention_mode spargeattn choice + --sparge_topk flag

Full changelog: md/changelog.md

v1.5.8 - TensorRT VAE Feather Weight Buffers Downgraded to FP16 (Encoder & Decoder)

Choose a tag to compare

@ussoewwin ussoewwin released this 02 Oct 04:20
EN 中文

TensorRT VAE — Feather Weight Buffers Downgraded to FP16 (Encoder & Decoder)


1. Overview of the VRAM Issue

Inside the single-shot tiled TensorRT VAE execution (_encode_single_chunk() / _decode_single_chunk()), two full-resolution accumulation buffers are allocated per chunk call:

Buffer Purpose Shape (decoder) Shape (encoder) Precision (before this fix)
result Accumulates every tile's output before the final normalization division (1, 3, video_frames, raw_out_h, raw_out_w) (1, 32, latent_frames, raw_h, raw_w) FP32
weights The feather (blend) map that steers that division identical shape identical shape FP32 ← problem

The result buffer must remain FP32 — it accumulates every tile output in full precision and is divided exactly once at the end. The weights buffer, however, only exists to steer that division. Despite that, it was also allocated in FP32, wasting exactly as much VRAM as the result buffer itself.

VRAM Consumption of the Feather Weights Buffer (Decoder, schematic)

 FP32 weights buffer (before)

 6 GiB ┤                                    ██
       │                                    ██
 4 GiB ┤                        ██          ██
       │           ██           ██          ██
 2 GiB ┤  ██       ██           ██          ██
       │  ██       ██           ██          ██
 0 GiB ┼──██───────██───────────██──────────██────
          73f      145f        217f        289f     (1080p / 1088p)

 FP16 weights buffer (after): exactly half of each bar

2. Why FP16 Is Safe for the Feather Weights

The feather map is not data. It is a smooth, deterministic per-pixel blend mask built by _feather() from torch.linspace ramps:

def _feather(length: int, overlap: int, left: bool, right: bool, device) -> torch.Tensor:
    weight = torch.ones(length, device=device, dtype=torch.float32)
    if left and overlap:
        weight[:overlap] = torch.linspace(0.0, 1.0, overlap + 1, device=device)[1:]
    if right and overlap:
        weight[-overlap:] = torch.minimum(
            weight[-overlap:], torch.linspace(1.0, 0.0, overlap + 1, device=device)[1:]
        )
    return weight

Its values live in [0, 1], are identical across all channels and frames, and contain no high-frequency content (the ramp changes smoothly over overlap pixels: 96 px for the encoder, 12 px latent / 96 px output for the decoder).

Consequently:

  • FP16 has ~3 decimal digits of mantissa precision in [0.5, 1] and even more relative precision near 0 (subnormal range extends to ~6e-8), so the worst-case relative representation error of a weight value is ~5e-4.
  • The weights buffer only appears in the final expression as a divisor: decoded = result / weights.clamp_min(1e-6). A relative error ε on the divisor shifts the output by at most ~ε — i.e. well under one 8-bit level of the final image.
  • result remains FP32, so the accumulation never loses precision; only the blend map is stored in half precision.

3. Code Changes

Target Files

# File Path Function Line(s)
1 src/core/trt_encoder.py _encode_single_chunk() L172, L210, L213
2 src/core/trt_decoder.py _decode_single_chunk() L179, L218, L221

Before (Problematic Implementation)

# Encoder (L172) / Decoder (L179) — both identical:
result = torch.zeros((...), device="cuda", dtype=torch.float32)
weights = torch.zeros_like(result)   # FP32: doubles the accumulator footprint
# Accumulation (encoder L210 / decoder L218):
weights[:, :, :, ly:ly + tile_lat, lx:lx + tile_lat] += window
#                                ^ window is FP32 here; FP32 += FP32
# Final normalization (encoder L213 / decoder L221):
encoded = (result / weights.clamp_min(1e-6))[...]
decoded = (result / weights.clamp_min(1e-6)).clamp(-2.0, 2.0)[...]

After (Optimized Implementation)

# 1) Buffer allocated in FP16 — exactly half the VRAM:
result = torch.zeros((...), device="cuda", dtype=torch.float32)   # stays FP32
weights = torch.zeros_like(result, dtype=torch.float16)           # now FP16

# 2) Accumulation: the FP32 window is cast on write (explicit dtype, no silent promote):
weights[:, :, :, ly:ly + tile_lat, lx:lx + tile_lat] += window.to(weights.dtype)

# 3) Normalization: clamp in FP16, then cast back up to FP32 for the division —
#    result stays FP32 end-to-end, only the divisor passes through FP16:
encoded = (result / weights.clamp_min(1e-6).to(result.dtype))[...]

Technical Rationale

Change Impact
weights = torch.zeros_like(result, dtype=torch.float16) Buffer footprint halved; result precision untouched.
window.to(weights.dtype) on accumulate Explicit narrowing at the write boundary; the smooth ramps of _feather() are exactly representable enough in FP16 (values in [0,1], monotone).
weights.clamp_min(1e-6).to(result.dtype) on divide The anti-zero clamp still runs (now in FP16, where 1e-6 is representable as a subnormal), and the divisor is promoted back to FP32 so result / weights executes in full precision.
FP16 subnormal range down to ~6e-8 The 1e-6 clamp floor is safely representable; no precision cliff near zero.

4. Measured VRAM Savings (Decoder Buffer)

weights shape = (1, 3, video_frames, raw_out_h, raw_out_w); FP32 → FP16 halves it exactly:

Scenario FP32 (before) FP16 (after) Saved
1080p / 73 frames 1.72 GiB 0.86 GiB ~0.85 GiB
1080p / 145 frames 3.41 GiB 1.71 GiB ~1.68 GiB
1088p / 289 frames 6.80 GiB 3.40 GiB ~3.37 GiB

Encoder-side savings are smaller ((1, 32, latent_frames, raw_h, raw_w) operates at 1/8 spatial resolution): ~0.04 GiB at 1080p/73f up to ~0.14 GiB at 1088p/289f — included in the fix at no cost.


5. Numerical Verification

  • PSNR of decoded output vs the FP32-feather reference: 84–85 dB at both 512 px and 1024 px.
  • 84 dB means the RMS pixel error is ~0.004% of full scale — roughly 4 orders of magnitude below any perceptible threshold (a 1-LSB change in an 8-bit pipeline corresponds to ~48 dB).
  • Stable across 15+ repeated runs — no drift, no accumulation effects (the weights buffer is rebuilt per chunk call and never persists across calls).

6. Summary

Item Before After
Feather weights precision (enc + dec) FP32 FP16
result accumulation precision FP32 FP32 (unchanged)
Decoder weights VRAM @ 1080p/289f 6.80 GiB 3.40 GiB (-3.37 GiB)
Output quality vs FP32 feather baseline PSNR 84–85 dB (indistinguishable)
Scope — trt_encoder.py + trt_decoder.py only; no other path touched

Full changelog: md/changelog.md

v1.5.7 - DiT VRAM Spike Stabilization — Complete Technical Guide

Choose a tag to compare

@ussoewwin ussoewwin released this 30 Sep 09:34
EN 中文

1. Overview of the VRAM Spike Issue

With ConvRot INT8 and NVFP4 quantization, the resident VRAM footprint of the DiT model weights has been significantly reduced (~3.1 GB → ~1.6 GB for the 3B model).

However, during actual inference, VRAM allocation spiked instantaneously and repeatedly during DiT execution steps, leading to unstable memory behavior and potential out-of-memory (OOM) conditions.

VRAM Consumption (Schematic Diagram)

 12 GB ┤                    ▲ ← Instantaneous spike per block
       │                   ╱╲
 10 GB ┤             ▲    ╱  ╲    ▲
       │            ╱╲  ╱    ╲  ╱╲
  8 GB ┤      ▲   ╱  ╲╱      ╲╱  ╲
       │     ╱╲ ╱                  ╲
  6 GB ┤────╱──╳────────────────────╲──── ← Resident Weight Baseline
       │   ╱
  4 GB ┤──╱
       └──────────────────────────────────
         Block 0  Block 8  Block 16  Block 31

These spikes were caused not by the weights, but by intermediate activation tensors generated during forward passes.
Even though weights are compressed via INT8/NVFP4, intermediate computation buffers are expanded in full FP16/BF16 precision. Thus, weight quantization alone cannot mitigate activation memory spikes.


2. Root Cause Analysis

Four distinct root causes were identified across the DiT execution pipeline:

Cause A: SDPA Attention Output Buffer Duplication (Primary Driver)

pytorch_varlen_attention() accumulated the attention output slices of individual windows into a Python list and concatenated them via torch.cat() at the end:

# List accumulation followed by full concatenation:
output_splits.append(output_i)
return torch.cat(output_splits, dim=0)

At the exact moment torch.cat() is executed, all individual slice tensors in the list and the newly allocated concatenated tensor exist simultaneously in VRAM, effectively doubling the memory required for the attention output.

Additionally, invoking cu_seqlens.long().cpu() across q, k, and v triggered CPU-GPU synchronization barriers per window, stalling the CUDA pipeline.

Impact: At 1080p with 5 frames ($L \approx 16,320$ tokens), this caused an instantaneous spike of ~640 MB. At 21 frames, the spike exceeded ~2 GB.

Cause B: SwiGLU MLP gate/up/hidden Simultaneous 3-Tensor Materialization

The SwiGLU MLP in the 3B model expands the intermediate feature dimension using expand_ratio=4:

# Although written as a single line, three large tensors coexist simultaneously:
x = self.proj_out(F.silu(self.proj_in_gate(x)) * self.proj_in(x))
#                       ^^^^^^^^^^^^^^^^^^^^^^^^   ^^^^^^^^^^^^^^
#                       gate: (L, 6912) BF16       up: (L, 6912) BF16
#                              ↓ silu(gate) * up = hidden: (L, 6912) BF16

At $L = 50,000$ tokens, a single tensor consumes ~660 MB. Having gate, up, and hidden live at the same instant consumes ~2.0 GB per block.
With 32 transformer blocks executing sequentially, this spike recurred on every single block.

Cause C: Text Token Replication via Fancy Indexing

In Swin Window Attention, text tokens are replicated across all spatial-temporal windows using torch.cat([vid, txt])[tgt_idx].

Under the hood, this requires two stages:

  1. torch.cat allocates an intermediate concatenated tensor.
  2. [tgt_idx] advanced indexing creates a second copy for gathering.

Because this was performed separately for Q, K, and V across both positive and negative CFG branches, intermediate buffers multiplied up to 6 times per block.

Cause D: Repeated Allocation of Condition Concatenation in Euler Sampling

Inside the Euler sampler loop:

vid = torch.cat([args.x_t, latents_cond], dim=-1)  # 16ch + 17ch = 33ch

This concatenation ran for both positive and negative CFG branches on every sampling step.
A new $(L, 33)$ tensor was allocated on each step while the prior step's tensor awaited garbage collection, leading to memory fragmentation in the caching allocator.


3. Modified Files

# File Path Modification Summary
P0-a src/models/dit_3b/attention.py SDPA output buffer pre-allocation & direct slice assignment
P0-b src/models/dit_7b/attention.py Same as above (7B architecture)
P1-a src/models/dit_3b/mlp.py SwiGLU MLP chunked execution ($L &gt; 8192$)
P1-b src/models/dit_7b/mlp.py Same as above (7B architecture)
P2-a src/models/dit_3b/na.py Replaced fancy indexing with torch.index_select
P2-b src/models/dit_7b/na.py Same as above (7B architecture)
P3 src/core/infer.py Euler condition buffer pre-allocation & in-place reuse

4. Code Diffs & Technical Rationale


P0: SDPA Attention — Pre-allocated Output Buffer & Direct Slice Writes

Target Files: src/models/dit_3b/attention.py / src/models/dit_7b/attention.py
Target Function: pytorch_varlen_attention()

Before (Problematic Implementation)

def pytorch_varlen_attention(q, k, v, cu_seqlens_q, cu_seqlens_k,
                             max_seqlen_q=None, max_seqlen_k=None,
                             dropout_p=0.0, softmax_scale=None,
                             causal=False, deterministic=False):
    # Repeated CPU transfers triggering CUDA synchronizations
    q_splits = list(torch.tensor_split(q, cu_seqlens_q[1:-1].long().cpu(), dim=0))
    k_splits = list(torch.tensor_split(k, cu_seqlens_k[1:-1].long().cpu(), dim=0))
    v_splits = list(torch.tensor_split(v, cu_seqlens_k[1:-1].long().cpu(), dim=0))

    # Slices accumulated into a Python list
    output_splits = []
    for q_i, k_i, v_i in zip(q_splits, k_splits, v_splits):
        q_i = q_i.permute(1, 0, 2).unsqueeze(0)
        k_i = k_i.permute(1, 0, 2).unsqueeze(0)
        v_i = v_i.permute(1, 0, 2).unsqueeze(0)

        output_i = F.scaled_dot_product_attention(
            q_i, k_i, v_i,
            dropout_p=dropout_p if not deterministic else 0.0,
            is_causal=causal
        )

        output_i = output_i.squeeze(0).permute(1, 0, 2)
        output_splits.append(output_i)  # All slice tensors kept alive simultaneously

    # 2x memory duplication during cat
    return torch.cat(output_splits, dim=0)

After (Optimized Implementation)

def pytorch_varlen_attention(q, k, v, cu_seqlens_q, cu_seqlens_k,
                             max_seqlen_q=None, max_seqlen_k=None,
                             dropout_p=0.0, softmax_scale=None,
                             causal=False, deterministic=False):
    total_len, num_heads, head_dim = q.shape

    # Pre-allocate output buffer once. All slice outputs write directly here.
    output = torch.empty_like(q)

    # Convert cu_seqlens to CPU once (minimizing CPU-GPU synchronization)
    cu_q_cpu = cu_seqlens_q.long().cpu()
    cu_k_cpu = cu_seqlens_k.long().cpu()
    num_seqs = len(cu_q_cpu) - 1

    drop_p = dropout_p if not deterministic else 0.0

    for i in range(num_seqs):
        q_start, q_end = cu_q_cpu[i].item(), cu_q_cpu[i + 1].item()
        k_start, k_end = cu_k_cpu[i].item(), cu_k_cpu[i + 1].item()

        q_i = q[q_start:q_end].permute(1, 0, 2).unsqueeze(0)
        k_i = k[k_start:k_end].permute(1, 0, 2).unsqueeze(0)
        v_i = v[k_start:k_end].permute(1, 0, 2).unsqueeze(0)

        out_i = F.scaled_dot_product_attention(
            q_i, k_i, v_i,
            dropout_p=drop_p,
            is_causal=causal,
        )

        # Write directly into pre-allocated buffer slice; no intermediate list accumulation
        output[q_start:q_end] = out_i.squeeze(0).permute(1, 0, 2)

    return output

Technical Rationale

Change Impact
Pre-allocate output = torch.empty_like(q) Eliminates memory duplication caused by torch.cat on accumulated slices.
In-place slice write output[q_start:q_end] = ... Each slice is written immediately into the buffer; out_i is freed upon the next iteration.
Single CPU conversion for cu_seqlens Reduces CPU-GPU synchronization stalls from 3 per window to 1 for the entire function call.
Input slicing q[q_start:q_end] Uses zero-copy tensor views instead of torch.tensor_split list creation.

P1: SwiGLU MLP — Chunked Activation Forward

Target Files: src/models/dit_3b/mlp.py / src/models/dit_7b/mlp.py
Target Class: SwiGLUMLP

Before (Problematic Implementation)

class SwiGLUMLP(nn.Module):
    def __init__(self, dim, expand_ratio, multiple_of=256, operations=None):
        super().__init__()
        ops = operations if operations is not None else nn
        hidden_dim = int(2 * dim * expand_ratio / 3)
        hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)
        self.proj_in_gate = ops.Linear(dim, hidden_dim, bias=False)
        self.proj_out = ops.Linear(hidden_dim, dim, bias=False)
        self.proj_in = ops.Linear(dim, hidden_dim, bias=False)

    def forward(self, x):
        # Monolithic execution: gate, up, and hidden tensors coexist across all tokens L
        x = self.proj_out(F.silu(self.proj_in_gate(x)) * self.proj_in(x))
        return x

After (Optimized Implementation)

class SwiGLUMLP(nn.Module):
    # L <= 8192: Fast path (no chunking overhead)
    # L > 8192: Chunked path to limit activation VRAM spikes
    CHUNK_THRESHOLD = 8192

    def __init__(self, dim, ...
Read more

v1.5.6 - Loader-Only Model Downloads & INT8/NVFP4 Registry Updates

Choose a tag to compare

@ussoewwin ussoewwin released this 30 Sep 08:11
EN 中文

Part 1: Loader-Only Auto-Download (Download Policy Fix)

1. Pre-Fix Problems and Failure Modes

1.1 Forced Model Downloads on Every Install/Update

  • Problem: install.py ran a default-model pre-download step on every install/update (via scripts/download_models.py), unconditionally fetching seedvr2_ema_3b_fp8_e4m3fn.safetensors and ema_vae_fp16.safetensors.
  • Symptom: Users who neither want nor use those files had them re-downloaded on every single update — the only workaround was to keep the files forever, wasting disk space and bandwidth (#2).

1.2 Default DiT Fetch During TensorRT VAE Engine Build

  • Problem: When the loader-selected VAE file was missing, src/interfaces/trt_vae_model_loader.py called download_weight(DEFAULT_DIT, model, ...), which also pulled the default 3B FP8 DiT model.
  • Symptom: Building a TensorRT VAE engine could download an unrequested default DiT.

2. New Download Policy (Fixed)

Auto-download happens only when a model is selected in a loader node and the file is missing.

No other automatic downloads exist:

Trigger Download behavior
Installation / update None
Running the upscaler node Only the node-selected DiT / VAE models, if missing
TensorRT VAE engine build Only the loader-selected VAE, if missing
CLI (inference_cli.py) Only the models selected via CLI arguments

3. Fixed and Modified Files

File Status Change
install.py Modified Removed download_default_models() and its step-4 invocation; no pre-download on install/update
scripts/install.ps1 Modified Removed the forced default-model download block
src/interfaces/trt_vae_model_loader.py Modified Downloads only the loader-selected VAE (never the default DiT)
src/utils/downloads.py Modified download_weight() accepts dit_model / vae_model optionally; only explicitly passed names are processed (None = skip)

4. Key Code Changes

# install.py — step 4 policy change
# 4. Default models: intentionally NOT pre-downloaded.
#    Downloads happen only when a model is selected in a loader node and is missing.
# src/interfaces/trt_vae_model_loader.py — selected VAE only
if not (model_dir / model).exists():
    print(f"[SeedVR2 TensorRT] Downloading {model} to {model_dir}...")
    # Download only the model selected in the loader (never the default DiT).
    download_weight(vae_model=model, model_dir=str(model_dir))
# src/utils/downloads.py — optional slots
def download_weight(dit_model: Optional[str] = None, vae_model: Optional[str] = None, model_dir: Optional[str] = None, debug=None) -> bool:
    ...
    files_to_download = [
        (name, MODEL_REGISTRY.get(name))
        for name in (dit_model, vae_model)
        if name
    ]

Part 2: Six ConvRot INT8 / NVFP4 DiT Models Registered

Registered six quantized DiT packs in MODEL_REGISTRY (hosted on Comfy-Org/SeedVR2, diffusion_models/, SHA256-pinned). They are selectable and auto-downloadable under the policy above:

Model URL
seedvr2_3b_int8_convrot.safetensors https://huggingface.co/Comfy-Org/SeedVR2/resolve/main/diffusion_models/seedvr2_3b_int8_convrot.safetensors
seedvr2_3b_nvfp4.safetensors https://huggingface.co/Comfy-Org/SeedVR2/resolve/main/diffusion_models/seedvr2_3b_nvfp4.safetensors
seedvr2_7b_int8_convrot.safetensors https://huggingface.co/Comfy-Org/SeedVR2/resolve/main/diffusion_models/seedvr2_7b_int8_convrot.safetensors
seedvr2_7b_nvfp4.safetensors https://huggingface.co/Comfy-Org/SeedVR2/resolve/main/diffusion_models/seedvr2_7b_nvfp4.safetensors
seedvr2_7b_sharp_int8_convrot.safetensors https://huggingface.co/Comfy-Org/SeedVR2/resolve/main/diffusion_models/seedvr2_7b_sharp_int8_convrot.safetensors
seedvr2_7b_sharp_nvfp4.safetensors https://huggingface.co/Comfy-Org/SeedVR2/resolve/main/diffusion_models/seedvr2_7b_sharp_nvfp4.safetensors

Part 3: GGUF DiT Entries Removed

  • Removed the GGUF (Q4_K_M / Q8_0) DiT entries from MODEL_REGISTRY: seedvr2_ema_3b-Q4_K_M.gguf, seedvr2_ema_3b-Q8_0.gguf, seedvr2_ema_7b-Q4_K_M.gguf, seedvr2_ema_7b_sharp-Q4_K_M.gguf.

Full changelog: md/changelog.md

v1.5.5 - Three Core Improvements and Encoder Refactoring

Choose a tag to compare

@ussoewwin ussoewwin released this 28 Sep 10:24
EN 中文

Part 1: Three Core Improvements (Studio Production Engineering Knowledge)

1. Pre-Fix Problems and Failure Modes

1.1 Timestamp Drift & Playback Endpoint Freezes during Video Assembly

  • Problem: When long video renders are split into temporal chunks or batches and assembled back into an MP4 container, standard video encoders produce non-monotonic Presentation Time Stamps (PTS) or Variable Frame Rate (VFR) containers.
  • Symptom: In players like PotPlayer, the final video frame freezes on screen while audio continues playing for several seconds, or audio drifts out of sync by hundreds of milliseconds over extended playback.

1.2 Black/Corrupted Output Tiles from Mutable set_tensor_address ExecutionContext Overwrite

  • Problem: In TensorRT's execution model, an ExecutionContext holds mutable memory pointers assigned via set_tensor_address. Naive multi-tile or asynchronous CUDA stream dispatch triggers race conditions where a later queued tile overwrites the tensor memory address of an actively executing tile.
  • Symptom: Earlier tiles read or write to corrupted memory ranges, outputting black tiles, severe gray checkerboard artifacts, or completely blank outputs.

1.3 Progressive VRAM Fragmentation and Memory Leak in Multi-Batch Loops

  • Problem: Long upscale runs (50–200+ frames) generate large intermediate FP32 spatial accumulation buffers and latent chunks. When loops rely solely on standard Python garbage collection or simple del, circular references and PyTorch's caching allocator hold onto CUDA memory pool blocks without returning them to the OS.
  • Symptom: Available VRAM progressively degrades with each processed chunk, eventually triggering an Out-Of-Memory (CUDA out of memory) crash on long renders.

2. Newly Created and Modified Files

File Path Status Role in Theme 1
src/interfaces/video_save.py Newly Created Implements dedicated CFR video assembler node applying Studio's production FFmpeg flags (-fflags +genpts -avoid_negative_ts make_zero -fps_mode cfr).
src/core/trt_decoder.py Modified Adds _DECODE_LOCK mutual exclusion, synchronous 1-tile execution (stream.synchronize()), and per-chunk tri-partite memory cleanup.
src/core/trt_encoder.py Modified Adds _ENCODE_LOCK mutual exclusion, synchronous 1-tile execution (stream.synchronize()), and per-chunk tri-partite memory cleanup.
src/core/generation_phases.py Modified Applies deterministic tri-partite memory reclamation (del + gc.collect() + torch.cuda.empty_cache()) across Phase 1, Phase 2, and Phase 3.

3. Complete Source Code of Newly Created & Modified Components (Unabridged)

3.1 Newly Created File: src/interfaces/video_save.py (Complete Node Implementation)

"""
SeedVR2 Save Video Node
High-reliability video encoder and assembler for ComfyUI.
Applies Studio-grade timestamp rectification (-fflags +genpts -avoid_negative_ts make_zero -fps_mode cfr)
to eliminate audio drift and endpoint freeze on long upscale renders.
"""

from __future__ import annotations

import os
import shutil
import subprocess
import time
from pathlib import Path
from typing import Dict, Any, Optional

import torch
import cv2
from comfy_api.latest import io

try:
    import folder_paths
except ImportError:
    folder_paths = None


def _get_ffmpeg() -> str:
    found = shutil.which("ffmpeg")
    if found:
        return found
    candidate_dirs = [
        Path(r"C:\Program Files\ffmpeg\bin"),
        Path(r"C:\Program Files\ffmpeg"),
        Path(r"C:\Program Files (x86)\ffmpeg\bin"),
        Path(r"C:\ffmpeg\bin"),
        Path(r"D:\ffmpeg\bin"),
        Path(__file__).resolve().parents[2] / "bin" / "ffmpeg" / "bin",
        Path(__file__).resolve().parents[2] / "bin",
    ]
    for d in candidate_dirs:
        candidate_file = d / "ffmpeg.exe"
        if candidate_file.exists():
            return str(candidate_file)
    try:
        import imageio_ffmpeg
        exe = imageio_ffmpeg.get_ffmpeg_exe()
        if exe and Path(exe).exists():
            return str(exe)
    except Exception:
        pass
    return "ffmpeg"


class SeedVR2SaveVideo(io.ComfyNode):
    """
    SeedVR2 Save Video Node
    
    Encodes video frames to MP4 with Studio's timestamp rectification:
    -fflags +genpts -avoid_negative_ts make_zero -fps_mode cfr
    Eliminates audio drift and end-of-video playback stutter completely.
    """

    @classmethod
    def define_schema(cls) -> io.Schema:
        return io.Schema(
            node_id="SeedVR2SaveVideo",
            display_name="SeedVR2 Save Video (CFR & Sync Safe)",
            category="SEEDVR2",
            description=(
                "Save video frames to MP4 with Studio-grade timestamp rectification. "
                "Applies -fflags +genpts -avoid_negative_ts make_zero -fps_mode cfr "
                "to eliminate audio desynchronization and endpoint freeze."
            ),
            inputs=[
                io.Image.Input("images",
                    tooltip="Upscaled video frames [T, H, W, C] in range [0, 1]."
                ),
                io.Float.Input("fps",
                    default=24.0,
                    min=1.0,
                    max=240.0,
                    step=0.01,
                    tooltip="Output video frame rate (strictly locked to Constant Frame Rate / CFR)."
                ),
                io.String.Input("filename_prefix",
                    default="SeedVR2",
                    tooltip="Output filename prefix. Saved under ComfyUI output directory."
                ),
                io.Int.Input("crf",
                    default=18,
                    min=0,
                    max=51,
                    step=1,
                    tooltip="H.264 CRF quality level (lower = higher quality, default: 18)."
                ),
                io.Combo.Input("preset",
                    options=["ultrafast", "superfast", "veryfast", "faster", "fast", "medium", "slow", "slower", "veryslow"],
                    default="medium",
                    tooltip="FFmpeg H.264 encoding preset (default: medium)."
                ),
                io.String.Input("audio_source_path",
                    default="",
                    optional=True,
                    tooltip="Optional path to source video or audio file to mux into the output video with timestamp alignment."
                ),
            ],
            outputs=[
                io.String.Output("video_path",
                    tooltip="Full path to the saved MP4 video file."
                )
            ]
        )

    @classmethod
    def execute(
        cls,
        images: torch.Tensor,
        fps: float = 24.0,
        filename_prefix: str = "SeedVR2",
        crf: int = 18,
        preset: str = "medium",
        audio_source_path: str = "",
    ) -> io.NodeOutput:
        if images.ndim != 4:
            raise ValueError(f"Expected 4D image tensor [T, H, W, C], got {tuple(images.shape)}")

        output_dir = folder_paths.get_output_directory() if folder_paths else "output"
        os.makedirs(output_dir, exist_ok=True)

        timestamp = time.strftime("%Y%m%d_%H%M%S")
        final_filename = f"{filename_prefix}_{timestamp}.mp4"
        final_path = Path(output_dir) / final_filename
        temp_path = final_path.with_name(f"{final_path.stem}_noaudio.mp4")

        T, H, W, C = images.shape
        cpu_frames = images.detach().cpu()
        if cpu_frames.dtype != torch.uint8:
            cpu_frames = (cpu_frames.clamp(0, 1) * 255.0).to(torch.uint8)

        frames_np = cpu_frames.numpy()
        del cpu_frames

        ffmpeg_bin = _get_ffmpeg()
        encode_command = [
            ffmpeg_bin, "-y", "-loglevel", "error",
            "-f", "rawvideo", "-pix_fmt", "rgb24" if C == 3 else "rgba",
            "-s:v", f"{W}x{H}", "-r", f"{fps:.9g}", "-i", "pipe:0",
            "-fflags", "+genpts", "-avoid_negative_ts", "make_zero",
            "-fps_mode", "cfr",
            "-an", "-c:v", "libx264", "-preset", preset, "-crf", str(crf),
            "-pix_fmt", "yuv420p", "-movflags", "+faststart", str(temp_path),
        ]

        encoder = subprocess.Popen(encode_command, stdin=subprocess.PIPE)
        write_error = None
        try:
            assert encoder.stdin is not None
            for i in range(T):
                frame = frames_np[i]
                if C == 4:
                    frame = cv2.cvtColor(frame, cv2.COLOR_RGBA2RGB)
                encoder.stdin.write(frame.tobytes())
        except BrokenPipeError as exc:
            write_error = exc
        finally:
            if encoder.stdin is not None:
                encoder.stdin.close()

        ret = encoder.wait()
        if ret != 0:
            temp_path.unlink(missing_ok=True)
            raise RuntimeError(f"FFmpeg video encoding failed with code {ret}")
        if write_error is not None:
            temp_path.unlink(missing_ok=True)
            raise write_error

        # Audio muxing with timestamp rectification
        has_audio = bool(audio_source_path and os.path.exists(audio_source_path))
        if has_audio:
            video_duration = T / max(fps, 1e-6)
            mux_cmd = [
                ffmpeg_bin, "-y", "-i", str(temp_path), "-i", str(audio_source_path),
                ...
Read more