#!/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 .sumsq float64[input_features], .count int64 scalar; optional .samples float32[rows,input_features] and .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()