Default video VAE decode to FP16

This commit is contained in:
Daniel Maddern 2026-08-14 19:50:36 +07:00
parent ea7712b4b2
commit a388aa64ff
8 changed files with 54 additions and 26 deletions

View file

@ -187,9 +187,9 @@ gate and is exact.
| H3 final layer | Yes | Bit-exact final AdaLN, target-row modulation, and video/audio patch rows; Comfy materializes the AdaLN and output-head biases through BF16 | | H3 final layer | Yes | Bit-exact final AdaLN, target-row modulation, and video/audio patch rows; Comfy materializes the AdaLN and output-head biases through BF16 |
| H3 DiT | Yes | Strict all-block numeric parity not achieved | | H3 DiT | Yes | Strict all-block numeric parity not achieved |
| Beta/RES sampler | Yes | H3 replay is bit-exact for all 12 captured sampler calls and RES update is bit-exact from captured denoised; direct raw-output to model-wrapper denoised conversion retains small BF16/state boundary deltas | | Beta/RES sampler | Yes | H3 replay is bit-exact for all 12 captured sampler calls and RES update is bit-exact from captured denoised; direct raw-output to model-wrapper denoised conversion retains small BF16/state boundary deltas |
| Video VAE decoder | Yes | Same final sampled latent decoded through direct VAE versus Comfy PNG frames: mean abs pixel delta `0.0833898`, max `6` | | Video VAE decoder | Yes | Direct VAE temporal assembly matches upstream after overlap fix. FP16 is the default Comfy-equivalent runtime path; cat benchmark VAE decode is `25.085s`. Use FP32 only for exact direct diagnostics. |
| Audio VAE/decode/mux | No | Preview intentionally produces video only | | Audio VAE/decode/mux | Yes | Direct decoder-only MiniMax H3 audio VAE returns stereo `32000 Hz` waveform and muxes with generated video. Native audio latent scaling is fixed. |
| End-to-end prompt-only FL2VA preview | Yes | Direct 320x192x22 dragon smoke videos generated with Sage2. Comfy-equivalent CPU joint AV seed initialization matches captured `initial.pt` exactly. Video-side wrapper denoised conversion is exact after BF16 raw-output materialization; remaining seed-exact final-latent delta is localized to regenerated text conditioning plus audio-side wrapper conversion | | End-to-end prompt-only FL2VA preview | Yes | Apples-to-apples warm cat benchmark is at Comfy parity: Comfy warm `150.26s`; direct warm after text conditioning `149.304s`; direct warm including text conditioning `151.465s`. |
| Full Ref2VA | No | References, vision conditioning, VAE encode, audio, and muxing are unimplemented | | Full Ref2VA | No | References, vision conditioning, VAE encode, audio, and muxing are unimplemented |
## Remaining Gates, In Dependency Order ## Remaining Gates, In Dependency Order
@ -197,14 +197,11 @@ gate and is exact.
Only these are outstanding. Do not recapture or revisit rows marked complete Only these are outstanding. Do not recapture or revisit rows marked complete
unless the checkpoint, Comfy version, prompt, or backend changes. unless the checkpoint, Comfy version, prompt, or backend changes.
1. **Final video latent layout.** Compare direct `unpatchify_video` output 1. **Feature/performance work.** Prompt-only FL2VA is now closed against the
with Comfy before entering the sampler. warm Comfy baseline. Optimize load/caching/sampling, then evaluate Sage3,
2. **Text/audio end-to-end closure.** Direct H3, video wrapper conversion, CUDA graphs, and multi-GPU execution.
RES update, VAE decode, and initial AV noise now match their captures. 2. **Full Ref2VA support.** Add references, vision conditioning, VAE encode,
Close regenerated text conditioning and audio-side wrapper conversion before and reference-path validation gates.
claiming identical seed output.
3. **Feature/performance work.** Only then add audio, Ref2VA/reference paths,
Sage3, CUDA graphs, and multi-GPU execution.
## Existing Tools And Their Intended Gate ## Existing Tools And Their Intended Gate

View file

@ -58,6 +58,8 @@ The direct runner must first match the model contract and output quality. Beatin
## 2026-08-13 VAE Debug Handoff ## 2026-08-13 VAE Debug Handoff
Resolved on `vae-decode-optimization`: direct VAE temporal overlap constants now match upstream, audio decode/mux is implemented, and Comfy-equivalent FP16 video VAE is the default runtime path. The 960x544x124 cat benchmark now matches warm Comfy performance: Comfy `150.26s`, direct `149.304s` after text conditioning, direct VAE decode `25.085s`.
Current saved latent and comparison assets live under: Current saved latent and comparison assets live under:
- Spark: `/home/daniel/StoryStudioAssets/H3-output/h3-blackwell-runtime` - Spark: `/home/daniel/StoryStudioAssets/H3-output/h3-blackwell-runtime`

View file

@ -45,5 +45,7 @@ Generation and latent-decode tools are quiet by default: they suppress ffmpeg ba
- `--profile-memory`: print memory checkpoints in `tools/direct_t2v_preview.py`. - `--profile-memory`: print memory checkpoints in `tools/direct_t2v_preview.py`.
- `--ffmpeg-loglevel info`: show ffmpeg details instead of the default `error` level. - `--ffmpeg-loglevel info`: show ffmpeg details instead of the default `error` level.
- `--quiet`: suppress JSON summary lines. - `--quiet`: suppress JSON summary lines.
- `--vae-dtype float16`: use Comfy-style FP16 video VAE decode in `tools/direct_t2v_preview.py` or `tools/decode_video_latent.py`; this is the default runtime path. Use `--vae-dtype float32` only for exact direct-path diagnostics. `tools/direct_t2v_preview.py` also accepts `H3_VAE_DTYPE`.
- `--vae-tile-size 256`: set the direct video VAE spatial tile size. `tools/direct_t2v_preview.py` also accepts `H3_VAE_TILE_SIZE`.
Standalone `tools/compare_*`, `tools/trace_*`, `tools/inspect_*`, and `tools/patch_comfy_*` scripts are debugging utilities and remain opt-in by being separate commands. Standalone `tools/compare_*`, `tools/trace_*`, `tools/inspect_*`, and `tools/patch_comfy_*` scripts are debugging utilities and remain opt-in by being separate commands.

View file

@ -18,5 +18,20 @@
"sampler": "res_multistep", "sampler": "res_multistep",
"seed": 440202, "seed": 440202,
"reference_comfy_sage3_seconds": 49.893, "reference_comfy_sage3_seconds": 49.893,
"measurement": "ComfyUI prompt execution time after warm-up" "measurement": "ComfyUI prompt execution time after warm-up",
"prompt_only_fl2va_cat_benchmark": {
"seed": 440407,
"prompt": "A playful orange tabby cat starts in an ordinary cozy living room in a normal house, afternoon light, sofa and rug. The cat crouches, jumps, and does one clean athletic backflip in slow motion. As the backflip completes there is a sharp cinematic cut: the cat lands perfectly on a glowing neon disco dance floor wearing oversized black sunglasses. Mirror ball reflections, colorful lights, joyful party energy, stylish and funny, clear before-and-after transformation.",
"comfy_warm_seconds": 150.26,
"direct_warm_after_text_conditioned_seconds": 149.304,
"direct_warm_including_text_conditioning_seconds": 151.465,
"direct_cold_through_audio_decode_seconds": 180.494,
"direct_video_vae_dtype": "float16",
"direct_video_vae_tile_size": 256,
"direct_video_vae_decode_seconds": 25.085,
"direct_output": "\\\\192.168.1.162\\StoryStudioAssets\\H3-output\\h3-blackwell-runtime\\direct-cat-benchmark-960x544-124f-seed440407-fp16-vae.mp4",
"direct_log": "\\\\192.168.1.162\\StoryStudioAssets\\H3-output\\h3-blackwell-runtime\\benchmarks\\direct-cat-benchmark-960x544-124f-seed440407-fp16-vae.log",
"comfy_output": "\\\\192.168.1.162\\StoryStudioAssets\\H3-output\\h3-blackwell-runtime\\comfy-cat-benchmark-960x544-124f-seed440407_00001_.mp4",
"comfy_metrics": "\\\\192.168.1.162\\StoryStudioAssets\\H3-output\\h3-blackwell-runtime\\benchmarks\\comfy-cat-benchmark-960x544-124f-seed440407.json"
}
} }

View file

@ -23,6 +23,14 @@ LATENTS_MEAN = (0.858090341091156, -0.9606591463088989, 1.0661640167236328, -0.5
LATENTS_STD = (1.2223774194717407, 1.2767263650894165, 1.68317747116088865, 1.7549455165863037, 1.5636216402053833, 2.194143533706665, 0.96531379222869875, 1.05698859691619875, 0.841948926448822, 0.7729952931404114, 1.8955937623977661, 0.946841835975647, 0.7996809482574463, 0.44988900423049925, 0.7197399735450745, 0.69362932443618775, 2.961095094680786, 2.7694199085235595, 3.0496184825897215, 2.1088054180145265, 3.276226282119751, 3.1627357006073, 2.28168129920959475, 2.6127843856811525) LATENTS_STD = (1.2223774194717407, 1.2767263650894165, 1.68317747116088865, 1.7549455165863037, 1.5636216402053833, 2.194143533706665, 0.96531379222869875, 1.05698859691619875, 0.841948926448822, 0.7729952931404114, 1.8955937623977661, 0.946841835975647, 0.7996809482574463, 0.44988900423049925, 0.7197399735450745, 0.69362932443618775, 2.961095094680786, 2.7694199085235595, 3.0496184825897215, 2.1088054180145265, 3.276226282119751, 3.1627357006073, 2.28168129920959475, 2.6127843856811525)
def dtype_from_name(name: str) -> torch.dtype:
return {
"float32": torch.float32,
"float16": torch.float16,
"bfloat16": torch.bfloat16,
}[name]
def _rms_norm(x: torch.Tensor, weight: torch.Tensor | None, eps: float) -> torch.Tensor: def _rms_norm(x: torch.Tensor, weight: torch.Tensor | None, eps: float) -> torch.Tensor:
if os.getenv("H3_VAE_FAST_OPS", "").lower() in {"1", "true", "yes", "on"}: if os.getenv("H3_VAE_FAST_OPS", "").lower() in {"1", "true", "yes", "on"}:
if weight is None: if weight is None:
@ -198,7 +206,7 @@ class MiniMaxH3VideoVAE(nn.Module):
self.register_buffer("pixel_std", torch.tensor(IMAGENET_STD, device=device).view(1, 3, 1, 1, 1), persistent=False) self.register_buffer("pixel_std", torch.tensor(IMAGENET_STD, device=device).view(1, 3, 1, 1, 1), persistent=False)
@classmethod @classmethod
def from_safetensors(cls, path: str | Path, *, device: str | torch.device = "cuda", tiling: bool = True, dtype: torch.dtype = torch.float32) -> "MiniMaxH3VideoVAE": def from_safetensors(cls, path: str | Path, *, device: str | torch.device = "cuda", tiling: bool = True, dtype: torch.dtype = torch.float16) -> "MiniMaxH3VideoVAE":
model = cls(device="meta", tiling=tiling) model = cls(device="meta", tiling=tiling)
expected = model.state_dict() expected = model.state_dict()
if os.getenv("H3_FAST_SAFETENSORS", "").lower() in {"1", "true", "yes", "on"}: if os.getenv("H3_FAST_SAFETENSORS", "").lower() in {"1", "true", "yes", "on"}:

View file

@ -10,6 +10,8 @@ from pathlib import Path
import torch import torch
from h3_blackwell_runtime.vae_decoder import dtype_from_name
def sync() -> None: def sync() -> None:
if torch.cuda.is_available(): if torch.cuda.is_available():
@ -26,14 +28,6 @@ def timed(stage: str, rows: list[dict], fn):
return value return value
def dtype_from_name(name: str) -> torch.dtype:
return {
"float32": torch.float32,
"float16": torch.float16,
"bfloat16": torch.bfloat16,
}[name]
def ffmpeg_command(loglevel: str, *parts: str) -> list[str]: def ffmpeg_command(loglevel: str, *parts: str) -> list[str]:
return ["ffmpeg", "-hide_banner", "-loglevel", loglevel, *parts] return ["ffmpeg", "-hide_banner", "-loglevel", loglevel, *parts]
@ -44,7 +38,7 @@ def main() -> None:
parser.add_argument("--output", type=Path) parser.add_argument("--output", type=Path)
parser.add_argument("--metrics", type=Path, required=True) parser.add_argument("--metrics", type=Path, required=True)
parser.add_argument("--implementation", choices=("direct", "upstream"), default="direct") parser.add_argument("--implementation", choices=("direct", "upstream"), default="direct")
parser.add_argument("--dtype", choices=("float32", "float16", "bfloat16"), default="float32") parser.add_argument("--dtype", choices=("float32", "float16", "bfloat16"), default="float16")
parser.add_argument("--no-tiling", action="store_true") parser.add_argument("--no-tiling", action="store_true")
parser.add_argument("--tile-size", type=int) parser.add_argument("--tile-size", type=int)
parser.add_argument("--tile-overlap", type=int) parser.add_argument("--tile-overlap", type=int)

View file

@ -6,7 +6,7 @@ from pathlib import Path
import torch import torch
from h3_blackwell_runtime.vae_decoder import MiniMaxH3VideoVAE from h3_blackwell_runtime.vae_decoder import MiniMaxH3VideoVAE, dtype_from_name
parser = argparse.ArgumentParser() parser = argparse.ArgumentParser()
@ -17,6 +17,8 @@ parser.add_argument("--implementation", choices=("direct", "upstream"), default=
parser.add_argument("--frames-dir", type=Path) parser.add_argument("--frames-dir", type=Path)
parser.add_argument("--ffmpeg-loglevel", default="error") parser.add_argument("--ffmpeg-loglevel", default="error")
parser.add_argument("--quiet", action="store_true") parser.add_argument("--quiet", action="store_true")
parser.add_argument("--vae-dtype", choices=("float32", "float16", "bfloat16"), default="float16")
parser.add_argument("--vae-tile-size", type=int, default=256)
args = parser.parse_args() args = parser.parse_args()
state = torch.load(args.latent, map_location="cuda", weights_only=False) state = torch.load(args.latent, map_location="cuda", weights_only=False)
@ -28,7 +30,9 @@ if args.implementation == "direct":
"/vae/minimax_h3_video_vae_fp16.safetensors", "/vae/minimax_h3_video_vae_fp16.safetensors",
device="cuda", device="cuda",
tiling=not args.no_tiling, tiling=not args.no_tiling,
dtype=dtype_from_name(args.vae_dtype),
).eval() ).eval()
vae.tile_size = args.vae_tile_size
else: else:
from safetensors.torch import load_file from safetensors.torch import load_file
@ -36,7 +40,10 @@ else:
vae = UpstreamMiniMaxH3VideoVAE(tiling=not args.no_tiling).to("cuda").eval() vae = UpstreamMiniMaxH3VideoVAE(tiling=not args.no_tiling).to("cuda").eval()
checkpoint = load_file("/vae/minimax_h3_video_vae_fp16.safetensors", device="cuda") checkpoint = load_file("/vae/minimax_h3_video_vae_fp16.safetensors", device="cuda")
checkpoint = {key: value.to(dtype=dtype_from_name(args.vae_dtype)) for key, value in checkpoint.items()}
missing, unexpected = vae.load_state_dict(checkpoint, strict=False) missing, unexpected = vae.load_state_dict(checkpoint, strict=False)
vae = vae.to(dtype=dtype_from_name(args.vae_dtype))
vae.tile_size = args.vae_tile_size
if not args.quiet: if not args.quiet:
print({"upstream_missing": len(missing), "upstream_unexpected": len(unexpected)}, flush=True) print({"upstream_missing": len(missing), "upstream_unexpected": len(unexpected)}, flush=True)
with torch.inference_mode(): with torch.inference_mode():
@ -61,4 +68,4 @@ subprocess.run([
], check=True) ], check=True)
raw.unlink() raw.unlink()
if not args.quiet: if not args.quiet:
print({"output": str(args.output), "frames": frames, "shape": tuple(pixels.shape), "tiling": not args.no_tiling, "implementation": args.implementation}) print({"output": str(args.output), "frames": frames, "shape": tuple(pixels.shape), "tiling": not args.no_tiling, "implementation": args.implementation, "vae_dtype": args.vae_dtype, "vae_tile_size": vae.tile_size})

View file

@ -20,7 +20,7 @@ from h3_blackwell_runtime.qwen3vl_text import Qwen3VLPromptConditioner
from h3_blackwell_runtime.sampler import sample_video_res_multistep from h3_blackwell_runtime.sampler import sample_video_res_multistep
from h3_blackwell_runtime.t2v import random_av_latents from h3_blackwell_runtime.t2v import random_av_latents
from h3_blackwell_runtime.token_refiner import H3TokenRefiner from h3_blackwell_runtime.token_refiner import H3TokenRefiner
from h3_blackwell_runtime.vae_decoder import MiniMaxH3VideoVAE from h3_blackwell_runtime.vae_decoder import MiniMaxH3VideoVAE, dtype_from_name
parser = argparse.ArgumentParser() parser = argparse.ArgumentParser()
@ -42,6 +42,8 @@ parser.add_argument("--save-audio-latent", type=Path)
parser.add_argument("--audio-output", type=Path) parser.add_argument("--audio-output", type=Path)
parser.add_argument("--mux-audio", action="store_true") parser.add_argument("--mux-audio", action="store_true")
parser.add_argument("--skip-decode", action="store_true") parser.add_argument("--skip-decode", action="store_true")
parser.add_argument("--vae-dtype", choices=("float32", "float16", "bfloat16"), default=os.getenv("H3_VAE_DTYPE", "float16"))
parser.add_argument("--vae-tile-size", type=int, default=int(os.getenv("H3_VAE_TILE_SIZE", "256")))
args = parser.parse_args() args = parser.parse_args()
started = time.perf_counter() started = time.perf_counter()
last_report = started last_report = started
@ -115,7 +117,8 @@ if args.save_audio_latent is not None:
report({"audio_latent": str(args.save_audio_latent)}) report({"audio_latent": str(args.save_audio_latent)})
if args.skip_decode: if args.skip_decode:
raise SystemExit(0) raise SystemExit(0)
vae = MiniMaxH3VideoVAE.from_safetensors("/vae/minimax_h3_video_vae_fp16.safetensors", device="cuda").eval() vae = MiniMaxH3VideoVAE.from_safetensors("/vae/minimax_h3_video_vae_fp16.safetensors", device="cuda", dtype=dtype_from_name(args.vae_dtype)).eval()
vae.tile_size = args.vae_tile_size
report_memory("vae_loaded") report_memory("vae_loaded")
with torch.inference_mode(): with torch.inference_mode():
pixels = vae.decode(latent.to(next(vae.parameters()).dtype))[:, :, :frames] pixels = vae.decode(latent.to(next(vae.parameters()).dtype))[:, :, :frames]