Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150 / collect_importance.py
Lottolabs's picture
Upload verified mixed BFP4/BFP8 checkpoint with MTP and evaluation evidence
12f320c verified
Raw History Blame Contribute Delete
11.7 kB
#!/usr/bin/env python3
"""Collect CPU reference linear-input diagonal second moments for TT PTQ ranking.
This is activation importance statistics, NOT a llama.cpp-compatible imatrix,
not weight optimization, and not QAT/QAD. Uses the original reference loader,
including restoration of all original FP32 GDN tensors. Processes one complete
bounded record at a time; never allocates or stores vocabulary logits. Hooks
accumulate sum(x**2) per input channel in float64 and count actual token rows.
"""
import argparse
import importlib.metadata
import json
import os
import time
from pathlib import Path
from reference_eval import REVISION, checkpoint_info, load_reference, records, sha256, write_json
def main():
parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
parser.add_argument("--model", type=Path, required=True)
parser.add_argument("--input", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True, help="New output directory")
parser.add_argument("--revision", default=REVISION)
parser.add_argument("--threads", type=int, default=16)
parser.add_argument("--interop-threads", type=int, default=1)
parser.add_argument("--max-tokens", type=int, default=2048, help="Reject longer records, never truncate")
parser.add_argument("--accumulation-chunk-rows", type=int, default=128)
parser.add_argument("--sample-rows", type=int, default=8, help="Uniformly spaced full-corpus input rows per module; 0 disables")
parser.add_argument("--checkpoint-every", type=int, default=32)
args = parser.parse_args()
if min(args.threads, args.interop_threads, args.accumulation_chunk_rows, args.checkpoint_every) < 1 or args.max_tokens < 2 or args.sample_rows < 0:
parser.error("Invalid positive thread/chunk/checkpoint/context count or negative sample rows")
os.environ["CUDA_VISIBLE_DEVICES"] = ""
os.environ["HF_HUB_OFFLINE"] = "1"
os.environ["TRANSFORMERS_OFFLINE"] = "1"
os.environ["OMP_NUM_THREADS"] = str(args.threads)
os.environ["MKL_NUM_THREADS"] = str(args.threads)
os.environ["TOKENIZERS_PARALLELISM"] = "false"
import numpy as np
import torch
torch.set_num_threads(args.threads)
torch.set_num_interop_threads(args.interop_threads)
torch.manual_seed(0)
torch.use_deterministic_algorithms(True)
started = time.perf_counter()
config, tensors, provenance = checkpoint_info(args.model, args.revision)
forbidden = {config[key] for key in ("image_token_id", "video_token_id", "vision_start_token_id", "vision_end_token_id") if key in config}
vocab_size = config["text_config"]["vocab_size"]
total_tokens = total_records = 0
for record in records(args.input, args.max_tokens, vocab_size, forbidden):
if record["split"] != "calibration":
raise ValueError("Importance must use calibration-only records, not validation/heldout")
total_tokens += len(record["token_ids"])
total_records += 1
if not total_records:
raise ValueError("Empty calibration corpus")
args.output.mkdir(parents=True, exist_ok=False)
metadata_path = args.output / "metadata.json"
metadata = {"schema_version": 1, "status": "loading", "kind": "reference_cpu_linear_input_diagonal_second_moments", "use": "TT precision sensitivity ranking; importance = sumsq/count. Not llama.cpp imatrix format; no training or quantized-weight optimization.", **provenance, "input_sha256": sha256(args.input), "script_sha256": sha256(Path(__file__)), "reference_loader_sha256": sha256(Path(__file__).with_name("reference_eval.py")), "arguments": {key: str(value) if isinstance(value, Path) else value for key, value in vars(args).items()}, "packages": {name: importlib.metadata.version(name) for name in ("torch", "transformers", "numpy", "safetensors", "accelerate")}, "total_records": total_records, "total_input_tokens": total_tokens, "completed_records": 0, "completed_input_tokens": 0, "statistic": "float64 sum of FP32 squared linear-input activations over all full-record token positions; count is unpadded token rows. Mean diagonal second moment = sumsq / count.", "memory_policy": "One batch-size-1 complete record <= max-tokens; CPU BF16 reference with original FP32 GDN weights; use_cache=False; no vocabulary projection/logits; reduction temp <= accumulation-chunk-rows * largest_linear_input_features FP32; persistent float64 channel sums and optional sample-rows FP32 activation rows per module.", "sampling": "Fixed approximately uniformly spaced token indices over the full deterministic record stream, not a first-record activation slice", "module_key_format": "Original checkpoint module name (strip .weight), keys <module>.sumsq float64[input_features], <module>.count int64 scalar; optional <module>.samples float32[rows,input_features] and <module>.sample_token_indices int64[rows]", "exclusions": "Only torch.nn.Linear input channels. Embedding lookup, GDN convolution, norms and A_log are not diagonal-linear importance entries; original reference weights remain unchanged."}
write_json(metadata_path, metadata)
handles = []
stats = {}
sample_positions = np.linspace(0, total_tokens - 1, min(args.sample_rows, total_tokens), dtype=np.int64) if args.sample_rows else np.empty(0, dtype=np.int64)
def save_statistics():
arrays = {}
sample_arrays = {}
module_metadata = {}
for name, state in stats.items():
arrays[name + ".sumsq"] = state["sumsq"].numpy()
arrays[name + ".count"] = np.asarray(state["count"], dtype=np.int64)
module_metadata[name] = {"input_features": state["sumsq"].numel(), "count": state["count"], "sample_rows": state["sample_count"]}
if args.sample_rows:
sample_arrays[name + ".samples"] = state["samples"][:state["sample_count"]].numpy()
sample_arrays[name + ".sample_token_indices"] = sample_positions[:state["sample_count"]]
for filename, payload in (("importance.npz", arrays), ("samples.npz", sample_arrays)):
if not payload:
continue
destination = args.output / filename
temporary = destination.with_suffix(".npz.tmp")
with temporary.open("wb") as handle:
np.savez(handle, **payload)
temporary.replace(destination)
metadata["modules"] = module_metadata
metadata["elapsed_seconds"] = time.perf_counter() - started
write_json(metadata_path, metadata)
def accumulate(name, inputs):
state = stats[name]
if inputs.device.type != "cpu" or inputs.shape[-1] != state["sumsq"].numel():
raise ValueError(f"Unexpected input shape/device: {name}")
# Model linear inputs are sequence-major batches. Flattening may produce
# at most one per-record copy if a module provides a noncontiguous view.
rows = inputs.detach().reshape(-1, inputs.shape[-1])
base = state["count"]
left, right = np.searchsorted(sample_positions, (base, base + rows.shape[0]))
if right > left:
indices = torch.from_numpy(sample_positions[left:right] - base)
state["samples"][left:right].copy_(rows.index_select(0, indices))
state["sample_count"] = int(right)
for start in range(0, rows.shape[0], args.accumulation_chunk_rows):
# clone avoids mutating original activations if an input is FP32.
block = rows[start:start + args.accumulation_chunk_rows].to(dtype=torch.float32, copy=True)
if not torch.isfinite(block).all():
raise ValueError(f"Nonfinite reference activation: {name}")
block.square_()
state["sumsq"].add_(block.sum(dim=0, dtype=torch.float64))
state["count"] += rows.shape[0]
try:
model, loading = load_reference(args.model, tensors, torch)
metadata.update(loading)
linears = {name: module for name, module in model.named_modules() if isinstance(module, torch.nn.Linear)}
if "lm_head" not in linears:
raise ValueError("Reference model has no linear lm_head")
for name, module in linears.items():
source_name = tensors[name + ".weight"]["checkpoint_key"].removesuffix(".weight")
stats[source_name] = {"sumsq": torch.zeros(module.in_features, dtype=torch.float64), "count": 0, "samples": torch.empty((len(sample_positions), module.in_features), dtype=torch.float32), "sample_count": 0}
if name != "lm_head":
def hook(_module, inputs, source_name=source_name):
accumulate(source_name, inputs[0])
handles.append(module.register_forward_pre_hook(hook))
metadata["status"] = "running"
write_json(metadata_path, metadata)
with torch.inference_mode():
for record in records(args.input, args.max_tokens, vocab_size, forbidden):
if record["split"] != "calibration":
raise ValueError("Input changed to include noncalibration record")
input_ids = torch.tensor([record["token_ids"]], dtype=torch.long, device="cpu")
hidden = model.model(input_ids=input_ids, use_cache=False, return_dict=True).last_hidden_state
# The lm_head input is the final normalized hidden state. Collect
# it directly instead of paying the 248320-column projection.
accumulate(tensors["lm_head.weight"]["checkpoint_key"].removesuffix(".weight"), hidden)
del hidden, input_ids
metadata["completed_records"] += 1
metadata["completed_input_tokens"] += len(record["token_ids"])
metadata["last_record_id"] = record["id"]
if metadata["completed_records"] % args.checkpoint_every == 0:
save_statistics()
print(json.dumps({key: metadata[key] for key in ("status", "completed_records", "completed_input_tokens", "elapsed_seconds")}), flush=True)
if sha256(args.input) != metadata["input_sha256"]:
raise ValueError("Input corpus changed during collection")
if metadata["completed_input_tokens"] != total_tokens or metadata["completed_records"] != total_records:
raise ValueError("Incomplete corpus traversal")
for name, state in stats.items():
if state["count"] != total_tokens or state["sample_count"] != len(sample_positions) or not torch.isfinite(state["sumsq"]).all():
raise ValueError(f"Incomplete or nonfinite module statistics: {name}")
metadata["status"] = "complete"
save_statistics()
metadata["importance_sha256"] = sha256(args.output / "importance.npz")
if args.sample_rows:
metadata["samples_sha256"] = sha256(args.output / "samples.npz")
write_json(metadata_path, metadata)
print(json.dumps({"status": "complete", "modules": len(stats), "records": total_records, "tokens": total_tokens, "elapsed_seconds": metadata["elapsed_seconds"]}), flush=True)
except BaseException as error:
# Keep only the last coherent completed-record checkpoint on failure;
# hooks may have partial current-record sums, so do not rewrite NPZ here.
metadata["status"] = "failed"
metadata["error"] = f"{type(error).__name__}: {error}"
metadata["elapsed_seconds"] = time.perf_counter() - started
write_json(metadata_path, metadata)
raise
finally:
for handle in handles:
handle.remove()
if __name__ == "__main__":
main()