#!/usr/bin/env python """Per-layer input rows, kept in three separate reservoirs by modality. `calib_sample.py` draws rows uniformly from the packed sequence. Uniform means *proportional*, and video is 98.6% of the rows, so the resulting `OutputsError` is a video error with a rounding error of text and audio mixed in. That is fine for choosing lambda -- it is the objective deepcompressor specifies -- but it makes the metric structurally unable to answer the question that matters here: what the rows lambda was *not* calibrated on end up paying. Same replay, same hooks, same quota logic as `calib_sample.py`; the only difference is that the reservoir is split by the transformer's own `video_indices` / `audio_indices` / `text_indices`, so each modality fills its own quota regardless of how few rows of it a step contains. Text is ~102 rows per step, which is why it needs its own reservoir rather than a mask applied afterwards: a proportional draw of 8192 rows yields about 110 text rows in total, and 110 rows do not estimate a per-channel error over 5376 channels. """ from __future__ import annotations import argparse 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 = ("text", "video", "audio") class ModalRowSampler: """One reservoir per modality, each filled to `quota` independently.""" __slots__ = ("quota", "per_call", "rows", "generator") def __init__(self, quota: int, num_calls: int, generator: torch.Generator) -> None: self.quota = quota self.per_call = max(1, -(-quota // max(num_calls, 1))) self.rows = {m: [] for m in MODALITIES} self.generator = generator @torch.no_grad() def update(self, x: torch.Tensor, masks: dict[str, torch.Tensor] | None) -> None: flat = x.reshape(-1, x.shape[-1]) parts = ({m: flat[masks[m]] for m in MODALITIES} if masks is not None else {"text": flat, "video": flat[:0], "audio": flat[:0]}) for m, part in parts.items(): if part.shape[0] == 0: continue if sum(r.shape[0] for r in self.rows[m]) >= self.quota: continue take = min(self.per_call, part.shape[0]) idx = torch.randint(0, part.shape[0], (take,), device=part.device, generator=self.generator) self.rows[m].append(part.index_select(0, idx).to(torch.float16).cpu()) def result(self) -> dict[str, torch.Tensor]: return {m: (torch.cat(r)[: self.quota] if r else torch.empty(0)) for m, r in self.rows.items()} def main() -> int: ap = argparse.ArgumentParser(description="Modality-split per-layer input samples") ap.add_argument("--caches", required=True) ap.add_argument("--out", required=True) ap.add_argument("--shard", required=True, help="i/n") ap.add_argument("--rows-per-layer", type=int, default=4096, help="TOTAL per modality") ap.add_argument("--model-path", default=None) ap.add_argument("--device", default="cuda:0") ap.add_argument("--attention-backend", default="_flash_3_hub") ap.add_argument("--max-steps", type=int, default=16) ap.add_argument("--seed", type=int, default=0) args = ap.parse_args() import bench from h3opt.svdquant_rules import target_linears from diffusers import ModularPipeline device = torch.device(args.device) shard_i, shard_n = (int(x) for x in args.shard.split("/")) t0 = time.perf_counter() pipe = ModularPipeline.from_pretrained(args.model_path or bench.DEFAULT_MODEL) pipe.load_components(names=["transformer"], dtype=torch.bfloat16) transformer = pipe.transformer.to(device).eval() if args.attention_backend: transformer.set_attention_backend(args.attention_backend) print(f"denoiser loaded in {time.perf_counter() - t0:.1f}s", flush=True) cache_dir = Path(args.caches) step_files = sorted(p for p in cache_dir.glob("*.pt") if ".cond" not in p.name)[shard_i::shard_n] if args.max_steps and len(step_files) > args.max_steps: stride = len(step_files) / args.max_steps step_files = [step_files[min(int(i * stride), len(step_files) - 1)] for i in range(args.max_steps)] quota = max(1, -(-args.rows_per_layer // shard_n)) print(f"shard {shard_i}/{shard_n}: {len(step_files)} steps, quota {quota} rows/modality/layer", flush=True) gen = torch.Generator(device=device).manual_seed(args.seed * 1000 + shard_i) targets = target_linears(transformer) samplers = {n: ModalRowSampler(quota, len(step_files), gen) for n in targets} state: dict = {"masks": None, "seq": 0} def set_layout(_m, _a, kwargs): pos = kwargs.get("position_ids") if pos is None: return None seq = int(pos.shape[0]) masks = {} for m, key in (("text", "text_indices"), ("video", "video_indices"), ("audio", "audio_indices")): v = torch.zeros(seq, dtype=torch.bool, device=device) idx = kwargs.get(key) if idx is not None: v[idx.to(device)] = True masks[m] = v state["masks"], state["seq"] = masks, seq return None h0 = transformer.register_forward_pre_hook(set_layout, with_kwargs=True) handles = [] for n, m in targets.items(): def hook(_m, inp, _n=n): x = inp[0] rows = x.reshape(-1, x.shape[-1]).shape[0] # The token refiner runs on the text rows alone, so its row count matches neither the # packed sequence nor a slice of it; those rows are text by construction. samplers[_n].update(x, state["masks"] if rows == state["seq"] else None) handles.append(m.register_forward_pre_hook(hook)) cond: dict[str, dict] = {} t1 = time.perf_counter() with torch.no_grad(): for k, path in enumerate(step_files): rec = torch.load(path, map_location="cpu", weights_only=False) clip = rec["clip"] if clip not in cond: cond.clear() cond[clip] = torch.load(cache_dir / f"{clip}.cond.pt", map_location="cpu", weights_only=False) kw = {kk: v for kk, v in rec.items() if kk not in ("outputs", "clip", "step") and torch.is_tensor(v)} kw.update(cond[clip]) kw = {kk: (v.to(device, torch.bfloat16) if v.is_floating_point() else v.to(device)) for kk, v in kw.items()} transformer(**kw, return_dict=False) print(f" {k+1}/{len(step_files)}", flush=True) h0.remove() for h in handles: h.remove() payload = {n: s.result() for n, s in samplers.items()} out = Path(args.out) out.parent.mkdir(parents=True, exist_ok=True) torch.save(payload, out) tot = sum(v.numel() * v.element_size() for d in payload.values() for v in d.values()) print(f"wrote {out} ({tot/1e9:.2f} GB, {len(payload)} layers) in " f"{time.perf_counter()-t1:.0f}s", flush=True) return 0 if __name__ == "__main__": raise SystemExit(main())