| |
| """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") |
|
|
| |
| |
| |
| |
| from calib_smooth_lowrank import quantize_nvfp4 |
|
|
|
|
| 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") |
| 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()) |
|
|