The Wan fp8 black frame was root-caused (scripts/fp8_layer_ablation.py, measured on B200 with the production torch._scaled_mm path): per-row fp8 scales each activation row by row_amax/448, and the text prompt is padded to 512 tokens (~all padding for a short prompt), so condition_embedder's text embedder divides a zero padding row by a zero scale, which infs and renders every frame black. That embedder's bias makes every downstream row non-zero, so the whole 30-block attn1/attn2/ffn stack is fp8-clean (fp8-except- condition_embedder measured cosine 0.9998 vs bf16, 0 non-finite; fp8- everywhere is 100% non-finite). So the blanket fp8 deny was heavier than needed for Wan. Remove fp8 from the Wan deny and keep only condition_embedder in bf16 via a new _FP8_FAMILY_EXCLUDE_NAME_TOKENS; auto now restores fp8 (the Blackwell ladder head) for Wan2.2-TI2V-5B and -T2V-A14B (shared DiT class and padded-text conditioning). Full-generation check (512x320, 25 frames, 30 steps, cache on and off): mixed-fp8 is non-black (mean luma 182.6 vs dense 181.2), more accurate than int8 (LPIPS 0.129 vs 0.180 no-cache, 0.224 vs 0.251 with FBCache), faster (49.9 vs 64.6 ms/step; int8 was a per-step regression vs the 59.8 ms/step dense), at the same memory (19.34 GB, both -20% vs dense). HunyuanVideo-1.5 keeps the fp8 deny: its MMDiT masks the padding text tokens to zero inside every block, so the per-block context stream (add_*_proj / to_add_out / ff_context) regenerates zero rows layer after layer (fp8 on only the main blocks is 100% non-finite) so no small exclude set exists and int8 stays. mxfp8 / nvfp4 remain denied for Wan (same per-row scaled_mm family, not separately validated). exclude_tokens_for_scheme now takes an optional family, threaded through the runtime quantiser and the offline prequant builder + validator so offline == runtime (a stale Wan fp8 checkpoint baked without the exclude is rejected and re-quantised rather than loaded). Adds scripts/fp8_layer_ablation.py (the per-layer ablation probe) and a mean-luma black-frame metric plus mixed-fp8 vs int8 configs to the video bench.
550 lines
24 KiB
Python
550 lines
24 KiB
Python
# SPDX-License-Identifier: AGPL-3.0-only
|
|
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
|
|
|
|
"""Per-layer fp8 ablation probe for video DiTs (Wan / Hunyuan): find which layers -- if any --
|
|
break under production per-row fp8 (torch._scaled_mm), or prove the failure is systemic.
|
|
|
|
Motivation: the dense video default auto-quantises the DiT; on Blackwell the auto ladder leads
|
|
with fp8. On Wan2.2 and HunyuanVideo-1.5 that renders every frame BLACK (int8 is clean); the
|
|
shipped fix denies fp8 for those families -> int8. The question this probe answers: is the black
|
|
frame caused by a small, nameable set of outlier layers (so we could keep fp8 on the rest, mixed),
|
|
or is it systemic to the _scaled_mm path (activation dynamic-quant overflow / accumulate) -- in
|
|
which case per-layer exclusion cannot help and int8 stays the right call?
|
|
|
|
Method (single-forward proxy; cheap): run one bf16 generation, capture a REAL mid-schedule DiT
|
|
input tuple via a forward_pre_hook, then compare the dense bf16 DiT output on that tuple to the
|
|
output after quantising layer SUBSETS with the production fp8 config. Because it reuses
|
|
``_make_quant_config`` / ``make_filter_fn`` from diffusion_transformer_quant, the GEMM path is
|
|
byte-identical to production (no MSLK).
|
|
|
|
Phase 0 (this file's default) is the gate:
|
|
0a proxy-validity : fp8-ALL must reproduce the black signal in one forward (cosine collapse or
|
|
non-finite output). If fp8-ALL is CLEAN in one forward, the black frame is
|
|
multi-step / cache compounding -> single-forward bisection cannot localise it
|
|
-> STOP, keep int8.
|
|
0b mechanism : (i) per-Linear finiteness hooks locate the first inf/NaN (a localisable
|
|
overflow layer); (ii) flip use_fast_accum True/False -- if it flips
|
|
black->clean the fix is a one-line accumulate flag, no per-layer work.
|
|
It also records per-Linear input outlier stats (within-token spread + amax) on the bf16 forward,
|
|
ready for the Phase 2 outlier ranking if a bucket localises.
|
|
|
|
Example:
|
|
CUDA_VISIBLE_DEVICES=1 python scripts/fp8_layer_ablation.py --family wan2.2-ti2v-5b --phase 0
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import copy
|
|
import gc
|
|
import json
|
|
import os
|
|
import sys
|
|
import time
|
|
from pathlib import Path
|
|
from typing import Any, Optional
|
|
|
|
os.environ.setdefault("BITSANDBYTES_NOWELCOME", "1")
|
|
os.environ.setdefault("HF_HOME", "/mnt/disks/unslothai/ubuntu/workspace_81/BACKUP_05/temp/hf_cache")
|
|
|
|
_REPO_ROOT = Path(__file__).resolve().parent.parent
|
|
_BACKEND_ROOT = _REPO_ROOT / "studio" / "backend"
|
|
for _p in (str(_BACKEND_ROOT), str(_REPO_ROOT / "scripts")):
|
|
if _p not in sys.path:
|
|
sys.path.insert(0, _p)
|
|
|
|
# Reuse the benchmark's real pipeline-build + family plumbing so the load is identical.
|
|
from video_speedmem_bench import _FAMILIES, PROMPT, _build_pipe, _import_diffusers # noqa: E402
|
|
from core.inference.diffusion_transformer_quant import ( # noqa: E402
|
|
DEFAULT_MIN_LINEAR_FEATURES,
|
|
TQ_FP8,
|
|
_REQUIRE_BF16_SCHEMES,
|
|
_make_quant_config,
|
|
make_filter_fn,
|
|
)
|
|
|
|
|
|
# ── cuda helpers ────────────────────────────────────────────────────────────────
|
|
def _empty() -> None:
|
|
import torch
|
|
|
|
gc.collect()
|
|
if torch.cuda.is_available():
|
|
torch.cuda.empty_cache()
|
|
|
|
|
|
def _detach_to_cpu(obj: Any) -> Any:
|
|
"""Recursively detach + move tensors in a nested args/kwargs structure to CPU."""
|
|
import torch
|
|
|
|
if torch.is_tensor(obj):
|
|
return obj.detach().to("cpu")
|
|
if isinstance(obj, tuple):
|
|
return tuple(_detach_to_cpu(o) for o in obj)
|
|
if isinstance(obj, list):
|
|
return [_detach_to_cpu(o) for o in obj]
|
|
if isinstance(obj, dict):
|
|
return {k: _detach_to_cpu(v) for k, v in obj.items()}
|
|
return obj
|
|
|
|
|
|
def _to_device(obj: Any, device: str) -> Any:
|
|
"""Recursively move tensors to ``device`` (dtypes preserved -- timestep stays int)."""
|
|
import torch
|
|
|
|
if torch.is_tensor(obj):
|
|
return obj.to(device)
|
|
if isinstance(obj, tuple):
|
|
return tuple(_to_device(o, device) for o in obj)
|
|
if isinstance(obj, list):
|
|
return [_to_device(o, device) for o in obj]
|
|
if isinstance(obj, dict):
|
|
return {k: _to_device(v, device) for k, v in obj.items()}
|
|
return obj
|
|
|
|
|
|
def _extract_tensor(out: Any):
|
|
import torch
|
|
|
|
if torch.is_tensor(out):
|
|
return out
|
|
s = getattr(out, "sample", None)
|
|
if torch.is_tensor(s):
|
|
return s
|
|
if isinstance(out, (list, tuple)) and out and torch.is_tensor(out[0]):
|
|
return out[0]
|
|
return out
|
|
|
|
|
|
# ── forward-tuple capture ────────────────────────────────────────────────────────
|
|
class _Stop(Exception):
|
|
pass
|
|
|
|
|
|
def _capture_forward_tuple(pipe, *, steps, width, height, num_frames, guidance, seed,
|
|
capture_call, gvg):
|
|
"""Run a bf16 generation and grab the ``capture_call``-th transformer forward's
|
|
(args, kwargs) to CPU, then abort. This is a REAL mid-schedule DiT input."""
|
|
import torch
|
|
|
|
holder: dict[str, Any] = {}
|
|
calls = [0]
|
|
|
|
def _pre(mod, args, kwargs):
|
|
calls[0] += 1
|
|
if calls[0] == capture_call:
|
|
holder["args"] = _detach_to_cpu(args)
|
|
holder["kwargs"] = _detach_to_cpu(kwargs)
|
|
raise _Stop()
|
|
return None
|
|
|
|
handle = pipe.transformer.register_forward_pre_hook(_pre, with_kwargs=True)
|
|
g = torch.Generator(device="cuda").manual_seed(seed)
|
|
kwargs = dict(
|
|
prompt=PROMPT, width=width, height=height, num_frames=num_frames,
|
|
num_inference_steps=steps, generator=g,
|
|
)
|
|
if gvg:
|
|
guider = getattr(pipe, "guider", None)
|
|
if guider is not None and hasattr(guider, "guidance_scale"):
|
|
try:
|
|
guider.guidance_scale = guidance
|
|
except Exception:
|
|
pass
|
|
else:
|
|
kwargs["guidance_scale"] = guidance
|
|
try:
|
|
pipe(**kwargs)
|
|
except _Stop:
|
|
pass
|
|
finally:
|
|
handle.remove()
|
|
if "kwargs" not in holder:
|
|
raise RuntimeError(
|
|
f"transformer forward never reached call #{capture_call} "
|
|
f"(saw {calls[0]}); lower --capture-call"
|
|
)
|
|
return holder
|
|
|
|
|
|
def _forward_out(module, tup: dict, device: str):
|
|
import torch
|
|
|
|
args = _to_device(tup.get("args", ()), device)
|
|
kwargs = _to_device(tup.get("kwargs", {}), device)
|
|
with torch.no_grad():
|
|
out = module(*args, **kwargs)
|
|
return _extract_tensor(out)
|
|
|
|
|
|
# ── scoring ──────────────────────────────────────────────────────────────────────
|
|
def _score(ref, cand) -> dict:
|
|
"""Compare a candidate DiT output to the bf16 reference. Non-finite entries are the
|
|
strongest black-frame signal; cosine / relL2 / norm-ratio are computed on finite entries."""
|
|
import torch
|
|
|
|
reff = ref.detach().flatten().float()
|
|
canf = cand.detach().flatten().float()
|
|
fin_c = torch.isfinite(canf)
|
|
frac_nonfinite = float(1.0 - fin_c.float().mean().item())
|
|
mask = fin_c & torch.isfinite(reff)
|
|
if mask.any():
|
|
a = canf[mask]
|
|
b = reff[mask]
|
|
cos = float(torch.nn.functional.cosine_similarity(a.unsqueeze(0), b.unsqueeze(0)).item())
|
|
rel = float((torch.linalg.vector_norm(a - b) / (torch.linalg.vector_norm(b) + 1e-12)).item())
|
|
norm_ratio = float((torch.linalg.vector_norm(a) / (torch.linalg.vector_norm(b) + 1e-12)).item())
|
|
else:
|
|
cos, rel, norm_ratio = 0.0, float("inf"), 0.0
|
|
return {
|
|
"cosine": round(cos, 5),
|
|
"relL2": round(rel, 5),
|
|
"frac_nonfinite": round(frac_nonfinite, 6),
|
|
"norm_ratio": round(norm_ratio, 5),
|
|
"cand_amax": round(float(canf[fin_c].abs().max().item()) if fin_c.any() else float("inf"), 3),
|
|
}
|
|
|
|
|
|
def _verdict(sc: dict) -> str:
|
|
"""Black-frame proxy verdict from a score dict."""
|
|
if sc["frac_nonfinite"] > 0:
|
|
return "BROKEN(non-finite)"
|
|
if sc["cosine"] < 0.5 or sc["norm_ratio"] < 0.3 or sc["norm_ratio"] > 3.0:
|
|
return "BROKEN(collapse)"
|
|
if sc["cosine"] < 0.98:
|
|
return "DEGRADED"
|
|
return "CLEAN"
|
|
|
|
|
|
# ── quant helpers ──────────────────────────────────────────────────────────────
|
|
def _fp8_filter(min_features: int, exclude_tokens=(), only_tokens=()):
|
|
"""A production fp8 filter (require_bf16) optionally narrowed:
|
|
- exclude_tokens: skip linears whose fqn contains any token (keep them bf16)
|
|
- only_tokens : quantise ONLY linears whose fqn contains any token."""
|
|
base = make_filter_fn(min_features, exclude_name_tokens=exclude_tokens,
|
|
require_bf16=(TQ_FP8 in _REQUIRE_BF16_SCHEMES))
|
|
|
|
def filt(module, fqn: str = "") -> bool:
|
|
if not base(module, fqn):
|
|
return False
|
|
name = (fqn or "").lower()
|
|
if exclude_tokens and any(t in name for t in exclude_tokens):
|
|
return False
|
|
if only_tokens and not any(t in name for t in only_tokens):
|
|
return False
|
|
return True
|
|
|
|
return filt
|
|
|
|
|
|
def _ablate(dense_cpu, tup, *, min_features, fast_accum=None, exclude_tokens=(), only_tokens=(),
|
|
instrument=False):
|
|
"""Deepcopy the dense DiT -> GPU -> fp8-quantise the selected subset -> one forward.
|
|
Returns (cpu_output_tensor, finiteness_records|None)."""
|
|
import torch
|
|
from torchao.quantization import quantize_
|
|
|
|
m = copy.deepcopy(dense_cpu).to("cuda")
|
|
filt = _fp8_filter(min_features, exclude_tokens=exclude_tokens, only_tokens=only_tokens)
|
|
quantize_(m, _make_quant_config(TQ_FP8, fast_accum=fast_accum), filter_fn=filt)
|
|
|
|
records = None
|
|
hooks = []
|
|
if instrument:
|
|
records = []
|
|
|
|
def mk(nm):
|
|
def hook(mod, inp, out):
|
|
o = out if torch.is_tensor(out) else (out[0] if isinstance(out, (list, tuple)) and out else None)
|
|
if o is None:
|
|
return
|
|
of = o.detach().float()
|
|
in_amax = None
|
|
in_frac_zero_rows = None
|
|
if inp and torch.is_tensor(inp[0]):
|
|
xin = inp[0].detach().float()
|
|
in_amax = float(xin.abs().max().item())
|
|
if xin.dim() >= 2:
|
|
ra = xin.reshape(-1, xin.shape[-1]).abs().amax(dim=-1)
|
|
in_frac_zero_rows = round(float((ra < 1e-6).float().mean().item()), 4)
|
|
records.append({
|
|
"name": nm,
|
|
"out_finite": bool(torch.isfinite(of).all().item()),
|
|
"out_amax": float(of.abs().max().item()) if torch.isfinite(of).any() else float("inf"),
|
|
"in_amax": in_amax,
|
|
"in_frac_zero_rows": in_frac_zero_rows,
|
|
})
|
|
return hook
|
|
|
|
for nm, mod in m.named_modules():
|
|
if isinstance(mod, torch.nn.Linear):
|
|
hooks.append(mod.register_forward_hook(mk(nm)))
|
|
|
|
out = _forward_out(m, tup, "cuda")
|
|
out = out.detach().to("cpu")
|
|
for h in hooks:
|
|
h.remove()
|
|
del m
|
|
_empty()
|
|
return out, records
|
|
|
|
|
|
# ── bf16 reference + input outlier stats (for Phase 2 ranking) ───────────────────
|
|
# ── Phase 1 buckets (fqn substring tokens), keyed by DiT kind ────────────────────
|
|
# Grounded in Phase 0: on Wan the inf originates in the layers whose ROW dimension is the
|
|
# padded text sequence (condition_embedder.text_embedder + cross-attn K/V that project
|
|
# encoder_hidden_states); the video self-attention (attn1) + attn2.to_q (video query) + FFN
|
|
# stay finite. "textpath" is the hypothesised minimal exclude set.
|
|
_BUCKETS: dict[str, dict[str, tuple[str, ...]]] = {
|
|
"wan": {
|
|
"textpath": ("text_embedder", "attn2.to_k", "attn2.to_v", "attn2.add_k", "attn2.add_v"),
|
|
"condition_embedder": ("condition_embedder",),
|
|
"attn2_all": ("attn2.",),
|
|
"attn2_kv": ("attn2.to_k", "attn2.to_v", "attn2.add_k", "attn2.add_v"),
|
|
"attn1": ("attn1.",),
|
|
"ffn": ("ffn.",),
|
|
},
|
|
"hunyuan": {
|
|
# candidate fix: the three input embedders consuming zero-padded / zero conditioning
|
|
# (image_embeds all-zero for T2V, encoder_hidden_states_2 ByT5 all-zero, pooled
|
|
# time_text_embed). "context_embedder" substring also matches "context_embedder_2".
|
|
"embedders": ("context_embedder", "image_embedder"),
|
|
"context_embedder": ("context_embedder",), # text refiner + context_embedder_2
|
|
"image_embedder": ("image_embedder",),
|
|
"context_embedder_2": ("context_embedder_2",),
|
|
"proj_out": ("proj_out",),
|
|
"main_blocks": ("transformer_blocks.",), # ONLY main blocks -> should be CLEAN
|
|
},
|
|
}
|
|
|
|
|
|
def _dit_kind(family: str) -> str:
|
|
return "hunyuan" if "hunyuan" in family.lower() else "wan"
|
|
|
|
|
|
def _zero_row_diag(tup: dict) -> dict:
|
|
"""Directly test the divide-by-zero hypothesis: fraction of near-zero (padding) rows in
|
|
the captured tensors. A per-row fp8 activation scale = row_amax / 448, so a row_amax of 0
|
|
yields scale 0 -> x/0 = inf. Padded text sequences (encoder_hidden_states) are the suspect."""
|
|
import torch
|
|
|
|
diag = {}
|
|
for key, val in tup.get("kwargs", {}).items():
|
|
if torch.is_tensor(val) and val.dim() >= 2 and val.is_floating_point():
|
|
x = val.float().reshape(-1, val.shape[-1])
|
|
row_amax = x.abs().amax(dim=-1)
|
|
n = int(row_amax.numel())
|
|
n_zero = int((row_amax < 1e-6).sum().item())
|
|
diag[key] = {
|
|
"rows": n,
|
|
"zero_amax_rows": n_zero,
|
|
"frac_zero_rows": round(n_zero / n, 4) if n else 0.0,
|
|
"min_row_amax": round(float(row_amax.min().item()), 8),
|
|
}
|
|
return diag
|
|
|
|
|
|
def _bf16_reference_and_stats(dense_gpu, tup, *, min_features):
|
|
"""Reference output + per-Linear input outlier stats on the bf16 forward. ``spread`` =
|
|
mean over tokens of (per-token channel amax / per-token channel median-abs): PR#150's
|
|
outlier proxy. Only linears the fp8 filter would quantise are recorded."""
|
|
import torch
|
|
|
|
keep = _fp8_filter(min_features) # which linears fp8 would touch
|
|
stats: list[dict] = []
|
|
hooks = []
|
|
|
|
def mk(nm):
|
|
def hook(mod, inp, out):
|
|
if not inp or not torch.is_tensor(inp[0]):
|
|
return
|
|
x = inp[0].detach().float()
|
|
x2 = x.reshape(-1, x.shape[-1]).abs() # [tokens, channels]
|
|
amax_tok = x2.amax(dim=-1)
|
|
med_tok = x2.median(dim=-1).values
|
|
spread = float((amax_tok / (med_tok + 1e-9)).mean().item())
|
|
stats.append({
|
|
"name": nm,
|
|
"in_amax": float(x2.max().item()),
|
|
"spread": round(spread, 2),
|
|
"in_features": int(mod.in_features),
|
|
"out_features": int(mod.out_features),
|
|
})
|
|
return hook
|
|
|
|
for nm, mod in dense_gpu.named_modules():
|
|
if isinstance(mod, torch.nn.Linear) and keep(mod, nm):
|
|
hooks.append(mod.register_forward_hook(mk(nm)))
|
|
|
|
ref = _forward_out(dense_gpu, tup, "cuda")
|
|
ref = ref.detach().to("cpu")
|
|
for h in hooks:
|
|
h.remove()
|
|
stats.sort(key=lambda s: s["spread"], reverse=True)
|
|
return ref, stats
|
|
|
|
|
|
def main(argv=None) -> int:
|
|
ap = argparse.ArgumentParser(description=__doc__)
|
|
ap.add_argument("--family", default="wan2.2-ti2v-5b", choices=sorted(_FAMILIES))
|
|
ap.add_argument("--phase", type=int, default=0, choices=[0, 1, 2])
|
|
ap.add_argument("--only", default="", help="phase 2: comma tokens -- fp8 ONLY these linears")
|
|
ap.add_argument("--exclude", default="", help="phase 2: comma tokens -- fp8 all EXCEPT these")
|
|
ap.add_argument("--steps", type=int, default=20)
|
|
ap.add_argument("--num-frames", type=int, default=25)
|
|
ap.add_argument("--width", type=int, default=512)
|
|
ap.add_argument("--height", type=int, default=320)
|
|
ap.add_argument("--seed", type=int, default=42)
|
|
ap.add_argument("--capture-call", type=int, default=8,
|
|
help="which transformer forward call to capture (~mid schedule)")
|
|
ap.add_argument("--min-features", type=int, default=DEFAULT_MIN_LINEAR_FEATURES)
|
|
ap.add_argument("--out", default="outputs/fp8_ablation")
|
|
args = ap.parse_args(argv)
|
|
|
|
import torch
|
|
|
|
out = Path(args.out)
|
|
out.mkdir(parents=True, exist_ok=True)
|
|
|
|
spec = _FAMILIES[args.family]
|
|
repo = spec["repo"]
|
|
force_fp32 = spec.get("vae_force_fp32", False)
|
|
guidance = spec.get("guidance", 5.0)
|
|
|
|
from core.inference.video_families import detect_video_family
|
|
|
|
fam_obj = detect_video_family(repo)
|
|
gvg = bool(getattr(fam_obj, "guidance_via_guider", False))
|
|
|
|
print(f"== fp8 layer ablation: family={args.family} repo={repo} phase={args.phase} ==", flush=True)
|
|
t0 = time.perf_counter()
|
|
pipe = _build_pipe(repo, force_fp32)
|
|
print(f"[load] pipe built in {time.perf_counter()-t0:.1f}s", flush=True)
|
|
|
|
tup = _capture_forward_tuple(
|
|
pipe, steps=args.steps, width=args.width, height=args.height,
|
|
num_frames=args.num_frames, guidance=guidance, seed=args.seed,
|
|
capture_call=args.capture_call, gvg=gvg,
|
|
)
|
|
ks = {k: (tuple(v.shape) if torch.is_tensor(v) else type(v).__name__)
|
|
for k, v in tup.get("kwargs", {}).items()}
|
|
print(f"[capture] call #{args.capture_call} kwargs={ks}", flush=True)
|
|
|
|
# Keep the dense bf16 DiT (still on GPU) as reference source; copy to CPU as the ablation seed.
|
|
dense_gpu = pipe.transformer
|
|
ref, stats = _bf16_reference_and_stats(dense_gpu, tup, min_features=args.min_features)
|
|
print(f"[bf16 ref] out shape={tuple(ref.shape)} amax={float(ref.float().abs().max()):.3f} "
|
|
f"n_quantized_linears={len(stats)}", flush=True)
|
|
|
|
dense_cpu = copy.deepcopy(dense_gpu).to("cpu")
|
|
# Free the pipeline (VAE/text encoders) to leave the GPU for one transformer at a time.
|
|
del pipe, dense_gpu
|
|
_empty()
|
|
|
|
report: dict[str, Any] = {"family": args.family, "repo": repo, "capture_call": args.capture_call,
|
|
"input_shapes": ks, "ref_amax": round(float(ref.float().abs().max()), 3),
|
|
"zero_row_diag": _zero_row_diag(tup)}
|
|
print(f"[zero-row diag] {report['zero_row_diag']}", flush=True)
|
|
|
|
if args.phase == 0:
|
|
# ── Phase 0a: proxy validity -- fp8-ALL (production accum = None auto-detect) ──
|
|
out_all, _ = _ablate(dense_cpu, tup, min_features=args.min_features, fast_accum=None)
|
|
sc_all = _score(ref, out_all)
|
|
report["fp8_all_autoaccum"] = {**sc_all, "verdict": _verdict(sc_all)}
|
|
print(f"[0a fp8-ALL auto-accum] {sc_all} -> {_verdict(sc_all)}", flush=True)
|
|
|
|
# ── Phase 0b-ii: accumulate flip ──
|
|
for fa in (False, True):
|
|
o, _ = _ablate(dense_cpu, tup, min_features=args.min_features, fast_accum=fa)
|
|
sc = _score(ref, o)
|
|
report[f"fp8_all_fast_accum_{fa}"] = {**sc, "verdict": _verdict(sc)}
|
|
print(f"[0b accum={fa}] {sc} -> {_verdict(sc)}", flush=True)
|
|
|
|
# ── Phase 0b-i: per-Linear finiteness (instrumented fp8-ALL forward) ──
|
|
_, records = _ablate(dense_cpu, tup, min_features=args.min_features, fast_accum=None,
|
|
instrument=True)
|
|
if records:
|
|
nonfinite = [r for r in records if not r["out_finite"]]
|
|
report["first_nonfinite_linears"] = nonfinite[:10]
|
|
report["n_nonfinite_linears"] = len(nonfinite)
|
|
finite_sorted = sorted([r for r in records if r["out_finite"]],
|
|
key=lambda r: r["out_amax"], reverse=True)
|
|
report["top_out_amax_linears"] = finite_sorted[:10]
|
|
print(f"[0b finiteness] {len(nonfinite)}/{len(records)} linears non-finite; "
|
|
f"first={[r['name'] for r in nonfinite[:5]]}", flush=True)
|
|
print(f"[0b top out_amax] {[(r['name'], round(r['out_amax'],1)) for r in finite_sorted[:5]]}",
|
|
flush=True)
|
|
|
|
report["top_spread_linears"] = stats[:20]
|
|
print(f"[stats] top-spread linears: "
|
|
f"{[(s['name'], s['spread']) for s in stats[:8]]}", flush=True)
|
|
|
|
va = report["fp8_all_autoaccum"]["verdict"]
|
|
if va.startswith("BROKEN"):
|
|
print("\n[GATE] fp8-ALL reproduces the black signal in one forward -> proxy VALID; "
|
|
"proceed to Phase 1 bucket ablation.", flush=True)
|
|
elif va == "CLEAN":
|
|
print("\n[GATE] fp8-ALL is CLEAN in one forward but full gen is black -> failure is "
|
|
"multi-step/cache; single-forward bisection cannot localise -> STOP, keep int8.",
|
|
flush=True)
|
|
else:
|
|
print(f"\n[GATE] fp8-ALL verdict={va}: borderline; inspect scores.", flush=True)
|
|
|
|
elif args.phase == 1: # ── Phase 1: bucket ablation (necessity + sufficiency) ──
|
|
kind = _dit_kind(args.family)
|
|
buckets = _BUCKETS[kind]
|
|
results: dict[str, Any] = {}
|
|
|
|
# baseline: fp8-ALL (should be BROKEN, matching Phase 0)
|
|
o, _ = _ablate(dense_cpu, tup, min_features=args.min_features)
|
|
results["fp8_all"] = {**_score(ref, o)}
|
|
results["fp8_all"]["verdict"] = _verdict(results["fp8_all"])
|
|
print(f"[1 baseline fp8-ALL] {results['fp8_all']['verdict']} {results['fp8_all']}", flush=True)
|
|
|
|
for bname, toks in buckets.items():
|
|
# NECESSITY: fp8 everything EXCEPT this bucket -- if CLEAN, this bucket is the culprit.
|
|
o, _ = _ablate(dense_cpu, tup, min_features=args.min_features, exclude_tokens=toks)
|
|
sc_ex = _score(ref, o)
|
|
# SUFFICIENCY: fp8 ONLY this bucket -- if BROKEN, this bucket alone reproduces damage.
|
|
o2, _ = _ablate(dense_cpu, tup, min_features=args.min_features, only_tokens=toks)
|
|
sc_only = _score(ref, o2)
|
|
results[bname] = {
|
|
"exclude": {**sc_ex, "verdict": _verdict(sc_ex)},
|
|
"only": {**sc_only, "verdict": _verdict(sc_only)},
|
|
}
|
|
print(f"[1 {bname}] EXCLUDE->{_verdict(sc_ex)} (cos {sc_ex['cosine']}, "
|
|
f"nf {sc_ex['frac_nonfinite']}) | ONLY->{_verdict(sc_only)} "
|
|
f"(cos {sc_only['cosine']}, nf {sc_only['frac_nonfinite']})", flush=True)
|
|
|
|
report["phase1"] = results
|
|
|
|
else: # ── Phase 2: instrument a specific only/exclude set, find first overflow ──
|
|
only = tuple(t.strip() for t in args.only.split(",") if t.strip())
|
|
exclude = tuple(t.strip() for t in args.exclude.split(",") if t.strip())
|
|
o, records = _ablate(dense_cpu, tup, min_features=args.min_features,
|
|
only_tokens=only, exclude_tokens=exclude, instrument=True)
|
|
sc = _score(ref, o)
|
|
report["phase2"] = {"only": only, "exclude": exclude, **sc, "verdict": _verdict(sc)}
|
|
print(f"[2 only={only} exclude={exclude}] {sc} -> {_verdict(sc)}", flush=True)
|
|
if records:
|
|
# first non-finite linears in execution order, with their input amax: an in_amax of
|
|
# ~0 at the first-broken layer confirms a zero-amax (padding) input row -> scale 0 -> inf.
|
|
nonfinite = [r for r in records if not r["out_finite"]]
|
|
report["phase2_first_nonfinite"] = nonfinite[:15]
|
|
report["phase2_n_nonfinite"] = len(nonfinite)
|
|
print(f"[2 finiteness] {len(nonfinite)}/{len(records)} quantised-path linears non-finite",
|
|
flush=True)
|
|
for r in nonfinite[:8]:
|
|
print(f" first-inf {r['name']} in_amax={r['in_amax']} "
|
|
f"in_frac_zero_rows={r['in_frac_zero_rows']} out_amax={r['out_amax']}",
|
|
flush=True)
|
|
|
|
dest = out / f"phase{args.phase}_{args.family}.json"
|
|
with open(dest, "w", encoding="utf-8") as fh:
|
|
json.dump(report, fh, indent=2)
|
|
print(f"wrote {dest}", flush=True)
|
|
return 0
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|