minimax-h3-nvfp4-lambda-modality / scripts /calib_sample_modal.py
yitongl's picture
Upload folder using huggingface_hub
0692312 verified
Raw
History Blame Contribute Delete
7.32 kB
#!/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())