#!/usr/bin/env python """What each modality's rows actually pay under each candidate lambda, with the real quantizer. `diag_lambda_crossmodal.py` answers this from the statistics with a bits-lost proxy. This answers it with the same NVFP4 fake-quantizer the calibration scored on, on real sampled rows, and with the rank-32 branch included -- so the number is comparable to the `err_smooth_lowrank` column in the calibration artifacts rather than merely correlated with it. Why it is needed at all: the calibration's own error column is measured on rows drawn uniformly from the packed sequence, and video is 98.6% of those rows. That column therefore ranks `lambda=video` best -- which it is, for video -- while being structurally unable to show what the text and audio rows paid for it. Those rows are 1.4% of the sequence and every video row attends to them. Errors are relative L2 of the quantized layer output against bf16, per modality, using each artifact's own lambda and its own low-rank factors. """ from __future__ import annotations import argparse import json import sys import time from pathlib import Path import torch REPO = Path(__file__).resolve().parent.parent sys.path.insert(0, str(REPO / "scripts")) sys.path.insert(0, str(REPO / "src")) MODALITIES = ("video", "text", "audio") # The whole point of this file is that its numbers sit next to `err_smooth_lowrank` in the # calibration artifacts, so the quantizer has to be the identical object, not a second # implementation of the same spec -- the per-tensor global scale under the FP8 group scale is easy # to leave out of a reimplementation and shifts every number if you do. from calib_smooth_lowrank import quantize_nvfp4 # noqa: E402 def err(x: torch.Tensor, w: torch.Tensor, lam: torch.Tensor, l1: torch.Tensor | None, l2: torch.Tensor | None) -> float: """Relative L2 of the SVDQuant forward against bf16, on these rows.""" ref = x @ w.T xs = x / lam ws = w * lam if l1 is not None and l1.abs().max() > 0: ws = ws - l1 @ l2 low = (xs @ l2.T) @ l1.T else: low = 0.0 y = quantize_nvfp4(xs) @ quantize_nvfp4(ws).T + low return float((y - ref).norm() / ref.norm().clamp_min(1e-12)) def main() -> int: ap = argparse.ArgumentParser() ap.add_argument("--samples", required=True, help="directory of msample_*.pt") ap.add_argument("--calibs", nargs="+", required=True, help="tag=path.pt") ap.add_argument("--out", required=True) ap.add_argument("--model-path", default=None) ap.add_argument("--shard", default="0/1", help="i/n over LAYERS") ap.add_argument("--max-rows", type=int, default=4096) ap.add_argument("--device", default="cuda:0") args = ap.parse_args() import bench from h3opt.svdquant_rules import target_linears from diffusers import ModularPipeline device = torch.device(args.device) cal = {} for spec in args.calibs: tag, path = spec.split("=", 1) cal[tag] = torch.load(path, map_location="cpu", weights_only=False)["layers"] tags = list(cal) i, n = (int(v) for v in args.shard.split("/")) names = sorted(next(iter(cal.values())))[i::n] pipe = ModularPipeline.from_pretrained(args.model_path or bench.DEFAULT_MODEL) pipe.load_components(names=["transformer"], dtype=torch.bfloat16) pipe.transformer.to("cpu") # one layer at a time on the GPU; 62 GB resident OOMs weights = {k: m.weight.detach().to("cpu") for k, m in target_linears(pipe.transformer).items() if k in names} del pipe torch.cuda.empty_cache() rows: dict[str, dict[str, list]] = {k: {m: [] for m in MODALITIES} for k in names} for p in sorted(Path(args.samples).glob("msample_*.pt")): blob = torch.load(p, map_location="cpu", weights_only=False) for k in names: if k not in blob: continue for m in MODALITIES: t = blob[k].get(m) if t is not None and t.numel(): rows[k][m].append(t) out: dict[str, dict] = {} t0 = time.perf_counter() for j, name in enumerate(names): w = weights[name].to(device, torch.float32) rec: dict[str, dict[str, float]] = {} for m in MODALITIES: if not rows[name][m]: continue x = torch.cat(rows[name][m])[: args.max_rows].to(device, torch.float32) rec[m] = {"rows": int(x.shape[0])} for t in tags: r = cal[t][name] rec[m][t] = err(x, w, r["lambda"].to(device, torch.float32), r["l1"].to(device, torch.float32), r["l2"].to(device, torch.float32)) del x out[name] = rec del w torch.cuda.empty_cache() if (j + 1) % 10 == 0 or j == len(names) - 1: print(f"[{j+1}/{len(names)}] {name} {rec.get('video', {})}", flush=True) p = Path(args.out) p.parent.mkdir(parents=True, exist_ok=True) p.write_text(json.dumps({"tags": tags, "shard": args.shard, "wall_s": round(time.perf_counter() - t0, 1), "layers": out}, indent=1)) print(f"wrote {p}", flush=True) return 0 if __name__ == "__main__": raise SystemExit(main())