Lottolabs's picture
Upload verified mixed BFP4/BFP8 checkpoint with MTP and evaluation evidence
12f320c verified
Raw History Blame Contribute Delete
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()