Download reference_eval.py from Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150: direct link, hf CLI and curl.
- Browser
- Download file 18.9 kB
-
https://huggingface.co/Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150/resolve/main/reference_eval.py
- Command line
-
hf download hf://Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150/reference_eval.py
-
curl -L -o reference_eval.py https://huggingface.co/Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150/resolve/main/reference_eval.py
18.9 kB
| #!/usr/bin/env python3 | |
| """Extract unquantized Qwen3.5 text logits with the native Transformers CPU model. | |
| Input JSONL: {"id": str, "split": str, "token_ids": [int, ...]}. | |
| Output records.jsonl contains those fields plus positions, logits_file and NLL. | |
| Each logits_file is a relative path to an uncompressed, C-order float32 NumPy | |
| .npy array [len(token_ids)-1, vocab_size]. Row p is the full-vocabulary logit | |
| vector after consuming token_ids[:p+1], predicting token_ids[p+1]. No prompt | |
| masking, BOS insertion, tokenization, padding, truncation or vocabulary reorder | |
| is performed. np.load(path, mmap_mode="r", allow_pickle=False) supports bounded | |
| KL(ref||candidate) evaluation; apply FP32 log_softmax to both raw logit arrays. | |
| """ | |
| import argparse | |
| import collections | |
| import hashlib | |
| import importlib.metadata | |
| import json | |
| import os | |
| from pathlib import Path | |
| import random | |
| import resource | |
| import struct | |
| import time | |
| REVISION = "c202236235762e1c871ad0ccb60c8ee5ba337b9a" | |
| TOKENIZER_FILES = ( | |
| "tokenizer.json", "tokenizer_config.json", "vocab.json", "merges.txt", | |
| "chat_template.jinja", | |
| ) | |
| def sha256(path): | |
| digest = hashlib.sha256() | |
| with path.open("rb") as handle: | |
| for block in iter(lambda: handle.read(1024 * 1024), b""): | |
| digest.update(block) | |
| return digest.hexdigest() | |
| def write_json(path, value): | |
| temporary = path.with_suffix(path.suffix + ".tmp") | |
| temporary.write_text(json.dumps(value, indent=2, sort_keys=True) + "\n") | |
| temporary.replace(path) | |
| def text_key(name): | |
| if name == "lm_head.weight": | |
| return name | |
| prefix = "model.language_model." | |
| if name.startswith(prefix): | |
| return "model." + name[len(prefix):] | |
| return None | |
| def checkpoint_info(model_dir, revision): | |
| config = json.loads((model_dir / "config.json").read_text()) | |
| if config.get("model_type") != "qwen3_5": | |
| raise ValueError("Expected the original Qwen3.5 multimodal checkpoint") | |
| if config.get("quantization_config"): | |
| raise ValueError("The reference checkpoint must be unquantized") | |
| index_path = model_dir / "model.safetensors.index.json" | |
| index = json.loads(index_path.read_text()) | |
| tensors = {} | |
| grouped_bytes = collections.Counter() | |
| for shard in sorted(set(index["weight_map"].values())): | |
| shard_path = model_dir / shard | |
| if shard_path.resolve().parent != model_dir.resolve(): | |
| raise ValueError(f"Shard is not a direct checkpoint file: {shard}") | |
| with shard_path.open("rb") as handle: | |
| header_size = struct.unpack("<Q", handle.read(8))[0] | |
| if header_size > 64 * 1024 * 1024: | |
| raise ValueError(f"Unreasonably large safetensors header: {shard}") | |
| header = json.loads(handle.read(header_size)) | |
| for name, value in header.items(): | |
| if name == "__metadata__": | |
| continue | |
| if index["weight_map"].get(name) != shard: | |
| raise ValueError(f"Tensor index/header mismatch: {name}") | |
| key = text_key(name) | |
| category = "text" if key else "vision" if name.startswith("model.visual.") else "mtp" if name.startswith("mtp.") else "other" | |
| grouped_bytes[category] += value["data_offsets"][1] - value["data_offsets"][0] | |
| if key: | |
| if value["dtype"] not in ("BF16", "F32"): | |
| raise ValueError(f"Unexpected reference tensor dtype: {name}: {value['dtype']}") | |
| tensors[key] = {"checkpoint_key": name, "shard": shard, **value} | |
| expected = {text_key(k) for k in index["weight_map"] if text_key(k)} | |
| if set(tensors) != expected or grouped_bytes["other"]: | |
| raise ValueError("Unexpected or missing checkpoint tensors") | |
| revisions = {} | |
| provenance_files = ["config.json", "model.safetensors.index.json", *TOKENIZER_FILES, *sorted(set(index["weight_map"].values()))] | |
| for name in provenance_files: | |
| sidecar = model_dir / ".cache/huggingface/download" / (name + ".metadata") | |
| if not sidecar.is_file(): | |
| raise ValueError(f"Missing upstream revision evidence: {sidecar}") | |
| observed = sidecar.read_text().splitlines()[0] | |
| if observed != revision: | |
| raise ValueError(f"Revision mismatch for {name}: {observed} != {revision}") | |
| revisions[name] = observed | |
| info = { | |
| "model_id": "Qwen/Qwen3.5-9B", | |
| "model_revision": revision, | |
| "model_path": str(model_dir.resolve()), | |
| "config_sha256": sha256(model_dir / "config.json"), | |
| "index_sha256": sha256(index_path), | |
| "revision_evidence": revisions, | |
| "weight_bytes_by_component": dict(grouped_bytes), | |
| "text_tensor_count": len(tensors), | |
| "tokenizer": { | |
| "model_id": "Qwen/Qwen3.5-9B", "revision": revision, | |
| "file_sha256": {name: sha256(model_dir / name) for name in TOKENIZER_FILES}, | |
| "applied_here": False, | |
| "vocabulary_order": "Original checkpoint lm_head row order, unchanged", | |
| }, | |
| "weight_integrity_note": "Revision sidecars plus config/index hashes; full multi-GB shard content hashes are not recomputed.", | |
| } | |
| return config, tensors, info | |
| def records(path, max_tokens, vocab_size, forbidden_tokens): | |
| seen = set() | |
| with path.open() as handle: | |
| for line_number, line in enumerate(handle, 1): | |
| if not line.strip(): | |
| continue | |
| record = json.loads(line) | |
| if not isinstance(record, dict): | |
| raise ValueError(f"Record {line_number} must be an object") | |
| identifier, split, ids = record.get("id"), record.get("split"), record.get("token_ids") | |
| if not isinstance(identifier, str) or not identifier or identifier in seen: | |
| raise ValueError(f"Record {line_number} has an empty/non-string/duplicate id") | |
| if not isinstance(split, str) or not split: | |
| raise ValueError(f"Record {identifier} requires a nonempty split") | |
| if not isinstance(ids, list) or not 2 <= len(ids) <= max_tokens: | |
| raise ValueError(f"Record {identifier} must have 2..{max_tokens} tokens; no truncation is allowed") | |
| if any(type(token) is not int or not 0 <= token < vocab_size for token in ids): | |
| raise ValueError(f"Record {identifier} has invalid token IDs") | |
| if forbidden_tokens.intersection(ids): | |
| raise ValueError(f"Record {identifier} contains image/video markers; this evaluator is text-only") | |
| seen.add(identifier) | |
| yield {"id": identifier, "split": split, "token_ids": ids} | |
| def load_reference(model_dir, tensors, torch): | |
| from safetensors import safe_open | |
| from transformers import Qwen3_5ForCausalLM, Qwen3_5TextConfig | |
| from transformers.models.qwen3_5 import modeling_qwen3_5 | |
| # These optional implementations are CUDA-only. Reject rather than silently | |
| # dispatch to a different implementation or mutate the serving environment. | |
| gpu_backends = ("FusedRMSNormGated", "causal_conv1d_fn", "causal_conv1d_update", "chunk_gated_delta_rule", "fused_recurrent_gated_delta_rule") | |
| active = [name for name in gpu_backends if getattr(modeling_qwen3_5, name, None) is not None] | |
| if active: | |
| raise RuntimeError(f"CPU reference requires the native PyTorch GDN fallback, but CUDA-only optional backends are installed: {active}") | |
| text_config = Qwen3_5TextConfig.from_pretrained(model_dir, local_files_only=True) | |
| text_config.use_cache = False | |
| # Native qwen3_5_text conversion renames model.language_model.* -> model.*. | |
| # The text-only class does not instantiate the visual encoder or MTP module. | |
| model, loading = Qwen3_5ForCausalLM.from_pretrained( | |
| model_dir, config=text_config, dtype=torch.bfloat16, | |
| device_map="cpu", local_files_only=True, use_safetensors=True, | |
| attn_implementation="sdpa", output_loading_info=True, | |
| ) | |
| for key in ("missing_keys", "mismatched_keys", "error_msgs"): | |
| if loading.get(key): | |
| raise RuntimeError(f"Incomplete reference load ({key}): {loading[key]}") | |
| unexpected = [key for key in loading.get("unexpected_keys", []) if not key.startswith(("mtp.", "model.visual."))] | |
| if unexpected: | |
| raise RuntimeError(f"Unexpected text checkpoint keys: {unexpected}") | |
| params = dict(model.named_parameters()) | |
| if set(params) != set(tensors): | |
| raise RuntimeError(f"Model/checkpoint parameter coverage mismatch: missing={set(tensors)-set(params)}, extra={set(params)-set(tensors)}") | |
| # dtype=BF16 otherwise downcasts the original checkpoint's small FP32 | |
| # A_log and GDN norm weights. Restore those original values, not a BF16 cast. | |
| fp32_by_shard = collections.defaultdict(list) | |
| for name, info in tensors.items(): | |
| if tuple(params[name].shape) != tuple(info["shape"]): | |
| raise RuntimeError(f"Shape mismatch: {name}") | |
| if info["dtype"] == "F32": | |
| fp32_by_shard[info["shard"]].append(name) | |
| for shard, names in fp32_by_shard.items(): | |
| with safe_open(model_dir / shard, framework="pt", device="cpu") as handle: | |
| for name in names: | |
| parent_name, leaf = name.rsplit(".", 1) | |
| original = handle.get_tensor(tensors[name]["checkpoint_key"]) | |
| setattr(model.get_submodule(parent_name), leaf, torch.nn.Parameter(original, requires_grad=False)) | |
| model.eval().requires_grad_(False) | |
| dtypes = collections.Counter() | |
| for name, parameter in model.named_parameters(): | |
| expected_dtype = torch.float32 if tensors[name]["dtype"] == "F32" else torch.bfloat16 | |
| if parameter.device.type != "cpu" or parameter.dtype != expected_dtype: | |
| raise RuntimeError(f"Unexpected parameter device/dtype: {name}: {parameter.device}/{parameter.dtype}") | |
| dtypes[str(parameter.dtype)] += parameter.numel() | |
| loading = {key: sorted(value) if isinstance(value, set) else value for key, value in loading.items()} | |
| return model, {"loading_info": loading, "parameter_elements_by_dtype": dict(dtypes), "native_fp32_tensors_restored": sum(map(len, fp32_by_shard.values()))} | |
| def main(): | |
| parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) | |
| parser.add_argument("--model", type=Path, required=True, help="Original unquantized checkpoint directory, mounted read-only") | |
| parser.add_argument("--input", type=Path, required=True, help="Pre-tokenized JSONL corpus") | |
| parser.add_argument("--output", type=Path, required=True, help="New output directory; must not already exist") | |
| 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("--seed", type=int, default=0) | |
| parser.add_argument("--dtype", choices=["bfloat16"], default="bfloat16") | |
| parser.add_argument("--max-tokens", type=int, default=512, help="Reject longer records rather than silently truncating") | |
| parser.add_argument("--logit-chunk-tokens", type=int, default=16, help="Bound temporary vocabulary projection memory") | |
| args = parser.parse_args() | |
| if min(args.threads, args.interop_threads, args.logit_chunk_tokens) < 1 or args.max_tokens < 2: | |
| parser.error("Thread/chunk counts must be positive and --max-tokens >= 2") | |
| 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(args.seed) | |
| np.random.seed(args.seed) | |
| random.seed(args.seed) | |
| torch.use_deterministic_algorithms(True) | |
| started = time.perf_counter() | |
| config, tensors, provenance = checkpoint_info(args.model, args.revision) | |
| vocab_size = config["text_config"]["vocab_size"] | |
| forbidden = {config[key] for key in ("image_token_id", "video_token_id", "vision_start_token_id", "vision_end_token_id") if key in config} | |
| record_counts = collections.Counter() | |
| token_counts = collections.Counter() | |
| for record in records(args.input, args.max_tokens, vocab_size, forbidden): | |
| record_counts[record["split"]] += 1 | |
| token_counts[record["split"]] += len(record["token_ids"]) - 1 | |
| if not record_counts: | |
| raise ValueError("Input corpus contains no records") | |
| args.output.mkdir(parents=True, exist_ok=False) | |
| (args.output / "logits").mkdir() | |
| metadata_path = args.output / "metadata.json" | |
| metadata = { | |
| "schema_version": 1, "status": "loading", "reference_kind": "unquantized_transformers_cpu", | |
| **provenance, "vocab_size": vocab_size, | |
| "input_sha256": sha256(args.input), "script_sha256": sha256(Path(__file__)), | |
| "packages": {name: importlib.metadata.version(name) for name in ("torch", "transformers", "numpy", "safetensors", "accelerate", "huggingface-hub")}, | |
| "settings": { | |
| "device": "cpu", "dtype": args.dtype, "preserve_checkpoint_fp32": True, | |
| "seed": args.seed, "threads": args.threads, "interop_threads": args.interop_threads, | |
| "deterministic_algorithms": True, "attention_implementation": "sdpa", | |
| "linear_attention_implementation": "native_pytorch_fallback", | |
| "use_cache": False, "mtp": False, "vision": False, "max_tokens": args.max_tokens, | |
| "logit_chunk_tokens": args.logit_chunk_tokens, "inference_mode": True, | |
| "teacher_forcing": "Full-context causal forward of identical input token_ids[:-1]; independent state per record", | |
| }, | |
| "array_contract": { | |
| "format": "npy", "dtype": "float32", "contents": "raw_logits", "shape": "[len(token_ids)-1, vocab_size]", | |
| "row_alignment": "row p consumes token_ids[:p+1] and predicts token_ids[p+1]", | |
| "logits_file_base": "output directory", "vocab_order": "checkpoint lm_head rows, unchanged", | |
| "nll": "sum/mean over every target token token_ids[1:], natural logarithms, FP32 logsumexp; no prompt masking", | |
| }, | |
| "record_counts_by_split": dict(record_counts), "target_tokens_by_split": dict(token_counts), | |
| "memory_estimate": { | |
| "resident_parameter_bytes": provenance["weight_bytes_by_component"]["text"], | |
| "logit_projection_temporary_bytes_approx": args.logit_chunk_tokens * vocab_size * 6, | |
| "largest_output_array_bytes": (args.max_tokens - 1) * vocab_size * 4, | |
| "note": "One record at a time; output is memory mapped; max_tokens bounds backbone activations. Peak RSS and mmap page cache are additional and are measured, not guaranteed by this estimate.", | |
| }, | |
| } | |
| write_json(metadata_path, metadata) | |
| totals = collections.defaultdict(lambda: {"nll_sum": 0.0, "nll_tokens": 0, "records": 0}) | |
| try: | |
| load_started = time.perf_counter() | |
| model, load_info = load_reference(args.model, tensors, torch) | |
| metadata.update(load_info) | |
| metadata["load_seconds"] = time.perf_counter() - load_started | |
| metadata["status"] = "running" | |
| write_json(metadata_path, metadata) | |
| with (args.output / "records.jsonl").open("x") as manifest, torch.inference_mode(): | |
| for ordinal, record in enumerate(records(args.input, args.max_tokens, vocab_size, forbidden)): | |
| record_started = time.perf_counter() | |
| ids = torch.tensor(record["token_ids"], dtype=torch.long, device="cpu") | |
| length = ids.numel() - 1 | |
| forward_started = time.perf_counter() | |
| hidden = model.model(input_ids=ids[:-1].unsqueeze(0), use_cache=False, return_dict=True, output_hidden_states=False, output_attentions=False).last_hidden_state | |
| backbone_seconds = time.perf_counter() - forward_started | |
| relative = f"logits/{ordinal:06d}.npy" | |
| destination = args.output / relative | |
| temporary = destination.with_suffix(".npy.partial") | |
| array = np.lib.format.open_memmap(temporary, mode="w+", dtype=np.float32, shape=(length, vocab_size)) | |
| nll_sum = 0.0 | |
| head_started = time.perf_counter() | |
| for start in range(0, length, args.logit_chunk_tokens): | |
| stop = min(start + args.logit_chunk_tokens, length) | |
| logits = model.lm_head(hidden[0, start:stop]).float() | |
| if not torch.isfinite(logits).all().item(): | |
| raise RuntimeError(f"Nonfinite reference logits: {record['id']} at position {start}") | |
| array[start:stop] = logits.numpy() | |
| target_logits = logits.gather(1, ids[start + 1:stop + 1, None]).squeeze(1) | |
| nll_sum += (torch.logsumexp(logits, dim=-1) - target_logits).double().sum().item() | |
| del logits, target_logits | |
| array.flush() | |
| del array, hidden, ids | |
| temporary.replace(destination) | |
| head_write_seconds = time.perf_counter() - head_started | |
| result = { | |
| **record, "positions": list(range(length)), "logits_file": relative, | |
| "nll_sum": nll_sum, "nll_mean": nll_sum / length, "nll_tokens": length, | |
| "timing": {"backbone_seconds": backbone_seconds, "head_and_write_seconds": head_write_seconds, "total_seconds": time.perf_counter() - record_started}, | |
| } | |
| manifest.write(json.dumps(result, separators=(",", ":")) + "\n") | |
| manifest.flush() | |
| total = totals[record["split"]] | |
| total["nll_sum"] += nll_sum | |
| total["nll_tokens"] += length | |
| total["records"] += 1 | |
| print(json.dumps({"id": record["id"], "split": record["split"], "nll_mean": result["nll_mean"], "seconds": result["timing"]["total_seconds"]}), flush=True) | |
| if sha256(args.input) != metadata["input_sha256"]: | |
| raise RuntimeError("Input corpus changed during evaluation") | |
| for total in totals.values(): | |
| total["nll_mean"] = total["nll_sum"] / total["nll_tokens"] | |
| metadata["nll_by_split"] = dict(totals) | |
| metadata["status"] = "complete" | |
| except BaseException as exc: | |
| metadata["status"] = "failed" | |
| metadata["error"] = f"{type(exc).__name__}: {exc}" | |
| raise | |
| finally: | |
| metadata["total_seconds"] = time.perf_counter() - started | |
| metadata["peak_rss_bytes"] = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * 1024 | |
| write_json(metadata_path, metadata) | |
| if __name__ == "__main__": | |
| main() | |