Download collect_importance.py from Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150: direct link, hf CLI and curl.
- Browser
- Download file 11.7 kB
-
https://huggingface.co/Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150/resolve/main/collect_importance.py
- Command line
-
hf download hf://Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150/collect_importance.py
-
curl -L -o collect_importance.py https://huggingface.co/Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150/resolve/main/collect_importance.py
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() | |