#!/usr/bin/env python3 """Teacher-forced, full-vocabulary TT evaluation; never run beside a live TT server. --describe validates inputs/precision and prints a launch description without importing torch/ttnn or opening hardware. Device runs require explicit ownership. Each N-token record produces N-1 rows: row p predicts token_ids[p+1]. """ import argparse import datetime import gc import hashlib import json import os from pathlib import Path import sys import time ROOT = Path(__file__).resolve().parent REVISION = "c202236235762e1c871ad0ccb60c8ee5ba337b9a" FAMILIES = {"attention", "gdn", "gdn_output", "mlp_gate_up", "mlp_down", "lm_head"} DTYPES = {"bf16", "bfp8", "bfp4"} PROFILES = {"baseline-bf16": "bf16", "baseline-bfp8": "bfp8", "current-bfp4": "bfp4", "serving-bfp4": "bfp4", "mixed-packed": "bfp4"} def digest(path): h = hashlib.sha256() with Path(path).open("rb") as stream: for block in iter(lambda: stream.read(1024 * 1024), b""): h.update(block) return h.hexdigest() def canonical(value): return json.dumps(value, sort_keys=True, separators=(",", ":")) def save_json(path, value): temporary = path.with_suffix(path.suffix + ".partial") temporary.write_text(json.dumps(value, indent=2, sort_keys=True) + "\n") temporary.replace(path) def precision_map(profile, overrides, layer_types): default = PROFILES[profile] result = {"lm_head": default, "layers": {}} for index, kind in enumerate(layer_types): attention = "attention" if kind == "full_attention" else "gdn" result["layers"][str(index)] = {attention: default, "mlp_gate_up": default, "mlp_down": default} if overrides is None: return result if profile == "serving-bfp4": raise ValueError("serving-bfp4 fixes the packed runtime and does not permit precision overrides") if not isinstance(overrides, dict) or set(overrides) - {"families", "layers"}: raise ValueError("Overrides must contain only families and/or layers objects") families = overrides.get("families", {}) if not isinstance(families, dict): raise ValueError("families must be an object") for family, dtype in families.items(): if family not in FAMILIES or dtype not in DTYPES: raise ValueError(f"Unsupported family/dtype: {family}={dtype}") if family == "lm_head": result[family] = dtype else: for values in result["layers"].values(): if family in values or (family == "gdn_output" and "gdn" in values): values[family] = dtype layers = overrides.get("layers", {}) if not isinstance(layers, dict): raise ValueError("layers must be an object keyed by canonical layer numbers") for index, values in layers.items(): if index not in result["layers"] or not isinstance(values, dict): raise ValueError(f"Unknown layer or non-object precision override: {index}") for family, dtype in values.items(): if (family not in result["layers"][index] and not (family == "gdn_output" and "gdn" in result["layers"][index])) or dtype not in DTYPES: raise ValueError(f"Unsupported layer family/dtype: {index}.{family}={dtype}") result["layers"][index][family] = dtype return result def load_records(path, vocab_size, context_limit): records, seen, split_tokens = [], set(), {} with path.open() as stream: for number, line in enumerate(stream, 1): if not line.strip(): continue row = json.loads(line) if not isinstance(row, dict) or not {"id", "split", "token_ids"} <= row.keys(): raise ValueError(f"Record {number}: expected id, split, token_ids") key, split, tokens = row["id"], row["split"], row["token_ids"] if isinstance(key, bool) or not isinstance(key, (str, int)) or canonical(key) in seen: raise ValueError(f"Record {number}: id must be a unique string/integer") if split not in {"calibration", "validation", "heldout"}: raise ValueError(f"Record {number}: split must be calibration, validation or heldout") if not isinstance(tokens, list) or not 2 <= len(tokens) <= context_limit: raise ValueError(f"Record {number}: token_ids length must be in [2,{context_limit}]") if any(type(token) is not int or not 0 <= token < vocab_size for token in tokens): raise ValueError(f"Record {number}: token ID outside vocabulary") token_hash = hashlib.sha256(canonical(tokens).encode()).hexdigest() if token_hash in split_tokens and split_tokens[token_hash] != split: raise ValueError("Identical token sequence appears in different splits") split_tokens[token_hash] = split seen.add(canonical(key)) records.append({"id": key, "split": split, "token_ids": tokens}) if not records: raise ValueError("Input dataset is empty") return records def runtime_environment(profile, inventory, mtp): if profile in ("serving-bfp4", "mixed-packed"): values = dict(inventory["serving_environment"]) else: values = { "MESH_DEVICE": "P150", "ARCH_NAME": "blackhole", "QWEN36_SINGLE_1D_DECODE": "1", "QWEN_GDN_RECURRENT_FUSED": "1", "QWEN_GDN_FRONTEND_FUSED": "1", } values.update({"QWEN36_MTP": "1" if mtp else "0", "QWEN_SDPA_BF8": "0"}) return values def load_text_weights(weights, torch, mtp): """Memory-map original text weights and optional lossless MTP; no HF allocation.""" from safetensors import safe_open from models.demos.blackhole.qwen36.tt.weight_mapping import remap_qwen36_state_dict if (weights / "native_manifest.json").is_file(): from native_checkpoint import load_native_checkpoint return load_native_checkpoint(weights, allow_unverified=True) index = json.loads((weights / "model.safetensors.index.json").read_text())["weight_map"] by_file = {} for name, filename in index.items(): if name.startswith("model.language_model.") or name == "lm_head.weight": by_file.setdefault(filename, []).append(name) if not by_file: raise ValueError("Expected original Qwen3.5 multimodal checkpoint text-weight names") raw = {} for filename, names in sorted(by_file.items()): shard = (weights / filename).resolve() if not shard.is_relative_to(weights): raise ValueError("Checkpoint index shard escapes read-only weights directory") with safe_open(str(shard), framework="pt", device="cpu") as source: for name in names: tensor = source.get_tensor(name) expected = torch.float32 if name.endswith((".linear_attn.A_log", ".linear_attn.norm.weight")) else torch.bfloat16 if tensor.dtype != expected: raise ValueError(f"Expected original {expected} weight, got {name}: {tensor.dtype}") raw[name] = tensor remapped = remap_qwen36_state_dict(raw) if not {"tok_embeddings.weight", "output.weight", "norm.weight"} <= remapped.keys(): raise ValueError("Text checkpoint is missing embedding/head/norm weights") if mtp: from models.demos.blackhole.qwen36.tt.weight_mapping import load_qwen36_mtp_state_dict remapped.update(load_qwen36_mtp_state_dict(weights)) return remapped def reset_sequence(model, kv_caches, kv_zero, native, ttnn): """Reset all causal history, preserving native GDN pool views when applicable.""" model.rope.rope_delta = 0 model._req_image_grid_thw = None model._req_video_grid_thw = None if model._last_hidden is not None: ttnn.deallocate(model._last_hidden) model._last_hidden = None if model.mtp is not None: model.mtp.reset_cache() for caches in kv_caches: for cache in caches: ttnn.copy(kv_zero, cache) if native: # This zeros the backing conv pool ONCE, not per-layer alias views. model._reset_dn_state_inplace() return for layer in model.layers: if layer.is_full_attention: layer.attention.reset_cache() continue dn = layer.attention # Generic TT recurrence consumes tiled state, not the native row-major # L1 pools allocated by this serving fork. Do not use model.reset_state: # it leaves use_inplace_state=True and drops external convolution views. dn.use_inplace_state = False dn._chunk_inplace_state = False dn.recurrent_state = ttnn.zeros( [1, dn.num_v_heads, dn.head_k_dim, dn.head_v_dim], dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=model.device, memory_config=ttnn.DRAM_MEMORY_CONFIG, ) dn.conv_state_q = dn.conv_state_k = dn.conv_state_v = None dn.fused_conv_state = None dn.split_conv_state = None gc.collect() def evaluate(args, records, metadata, cache, environment): # Scrub inherited serving flags BEFORE importing model/experimental modules. for name in list(os.environ): if name.startswith(("QWEN", "TT_QWEN")): del os.environ[name] os.environ.update(environment) os.environ.update({ "HF_MODEL": str(args.weights), "MODEL_WEIGHTS_DIR": str(args.weights), "TT_CACHE_PATH": str(cache), "HF_HUB_OFFLINE": "1", "TRANSFORMERS_OFFLINE": "1", "QWEN36_EVAL_PRECISION": canonical(metadata["precision"]), "TT_METAL_CACHE": str(cache / "kernel-cache"), }) import torch import numpy as np import ttnn from models.demos.blackhole.qwen36.tt.model import Qwen36Model from models.demos.blackhole.qwen36.tt.model_config import Qwen36ModelArgs imported_source = Path(sys.modules[Qwen36Model.__module__].__file__).resolve() expected_source = args.runtime.resolve() / "tt" / "model.py" if imported_source != expected_source: raise RuntimeError(f"Refusing non-isolated runtime: imported {imported_source}; expected {expected_source}") torch.set_num_threads(args.cpu_threads) mesh = ttnn.open_mesh_device( mesh_shape=ttnn.MeshShape(1, 1), l1_small_size=24576, num_command_queues=2, trace_region_size=0, ) try: if mesh.get_num_devices() != 1: raise RuntimeError("Evaluator supports one P150 only") model_args = Qwen36ModelArgs(mesh_device=mesh, max_batch_size=1, max_seq_len=metadata["max_seq_len"]) resolved_cache = model_args.weight_cache_path().resolve() if not resolved_cache.is_relative_to(cache): raise RuntimeError(f"Unsafe tensor cache path: {resolved_cache}") resolved_cache.mkdir(parents=True, exist_ok=True) state_dict = load_text_weights(args.weights, torch, args.mtp) model = Qwen36Model(mesh, model_args, state_dict, tensor_cache_path=resolved_cache) del state_dict gc.collect() if (model.mtp is not None) != args.mtp or model.vocab_size != metadata["vocab_size"]: raise RuntimeError("MTP/vocabulary contract violation") model._ondev_argmax = False native = environment.get("QWEN_GDN_RECURRENT_FUSED") == "1" blocks = metadata["max_seq_len"] // 64 kv_shape = [blocks, model_args.n_kv_heads, 64, model_args.head_dim] kv_caches = model.allocate_kv_caches(kv_shape, ttnn.bfloat16, batch_size=1) kv_zero = ttnn.zeros(kv_shape, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=mesh) page_table = ttnn.from_torch(torch.arange(blocks, dtype=torch.int32).reshape(1, -1), dtype=ttnn.int32, layout=ttnn.ROW_MAJOR_LAYOUT, device=mesh) if native: model._init_dn_zero_buffers() metadata.update({"torch_version": torch.__version__, "ttnn_version": getattr(ttnn, "__version__", None), "tensor_cache_path": str(resolved_cache), "imported_model_source": str(imported_source), "device_grid": str(mesh.compute_with_storage_grid_size())}) save_json(args.output / "metadata.json", metadata) generations = [] tokenizer = None if args.generate_tokens: from transformers import AutoTokenizer tokenizer = AutoTokenizer.from_pretrained(args.weights, local_files_only=True) with torch.inference_mode(), (args.output / "records.jsonl").open("x") as manifest: for record_index, record in enumerate(records): started = time.monotonic() reset_sequence(model, kv_caches, kv_zero, native, ttnn) prompt_token_ids = list(record["token_ids"]) token_ids = list(prompt_token_ids) count = len(token_ids) - 1 + args.generate_tokens filename = f"logits-{record_index:06d}.npy" target_path = args.output / filename temporary_path = args.output / (filename + ".partial") logits_array = np.lib.format.open_memmap(temporary_path, mode="w+", dtype=np.float32, shape=(count, model.vocab_size)) nll_sum = 0.0 token_seconds = [] for position in range(count): token = token_ids[position] token_started = time.monotonic() tokens_tt = ttnn.from_torch(torch.tensor([[token]], dtype=torch.int32), dtype=ttnn.uint32, layout=ttnn.ROW_MAJOR_LAYOUT, device=mesh) position_tt = ttnn.from_torch(torch.tensor([position], dtype=torch.int32), dtype=ttnn.int32, layout=ttnn.ROW_MAJOR_LAYOUT, device=mesh) cos, sin = model.rope.get_rot_mats(torch.tensor([[position]], dtype=torch.long)) output = model._forward_decode(tokens_tt, cos, sin, position_tt, page_table) host = ttnn.to_torch(output).float() if host.shape[-1] != model.vocab_size or host.numel() < model.vocab_size: raise RuntimeError(f"Invalid full-vocabulary logits shape: {tuple(host.shape)}") row = host.reshape(-1, model.vocab_size)[0] if not torch.isfinite(row).all().item(): raise RuntimeError(f"Non-finite logits at record {record['id']}, position {position}") logits_array[position] = row.numpy() stop_generation = False if args.generate_tokens and position >= len(prompt_token_ids) - 1: sampled = int(row.argmax()) token_ids.append(sampled) stop_generation = sampled == tokenizer.eos_token_id nll_sum += float(torch.logsumexp(row, dim=-1) - row[token_ids[position + 1]]) # RoPE slices can alias persistent tables: let references # release normally instead of force-deallocating their owner. for tensor in (output, tokens_tt, position_tt): ttnn.deallocate(tensor) token_seconds.append(time.monotonic() - token_started) print(canonical({"record": record["id"], "position": position, "token_elapsed_seconds": token_seconds[-1]}), flush=True) if stop_generation: count = position + 1 break logits_array.flush() if count < logits_array.shape[0]: trimmed = temporary_path.with_suffix(".trimmed") with trimmed.open("wb") as stream: np.save(stream, logits_array[:count]) trimmed.replace(temporary_path) del logits_array temporary_path.replace(target_path) item = {**record, "token_ids": token_ids, "positions": list(range(count)), "logits_file": filename, "logits_shape": [count, model.vocab_size], "logits_dtype": "float32", "nll_sum": nll_sum, "nll_mean": nll_sum / count, "nll_tokens": count, "token_elapsed_seconds": token_seconds, "decode_seconds_per_token": sum(token_seconds) / count, "elapsed_seconds": time.monotonic() - started, "logits_sha256": digest(target_path)} manifest.write(canonical(item) + "\n") manifest.flush() if args.generate_tokens: generated = token_ids[len(prompt_token_ids):] generations.append({"id": record["id"], "prompt_token_ids": prompt_token_ids, "generated_token_ids": generated, "text": tokenizer.decode(generated, skip_special_tokens=False), "seconds": time.monotonic() - started}) print(canonical({"record": record["id"], "split": record["split"], "nll_mean": item["nll_mean"]}), flush=True) if args.generate_tokens: (args.output / "generations.jsonl").write_text("".join(canonical(row) + "\n" for row in generations)) metadata["status"] = "complete" metadata["completed_records"] = len(records) finally: ttnn.close_mesh_device(mesh) def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--input", type=Path, required=True) parser.add_argument("--weights", type=Path, required=True) parser.add_argument("--output", type=Path, required=True) parser.add_argument("--cache-root", type=Path, default=ROOT / "candidate-cache") parser.add_argument("--runtime", type=Path, default=ROOT / "runtime" / "qwen36") parser.add_argument("--profile", choices=sorted(PROFILES), default="baseline-bf16") parser.add_argument("--overrides", type=Path) parser.add_argument("--context-limit", type=int, default=8192) parser.add_argument("--cpu-threads", type=int, default=8) parser.add_argument("--mtp", action="store_true", help="Load the MTP head while evaluating full target logits; does not run speculative generation") parser.add_argument("--generate-tokens", type=int, default=0, help="Free greedy generation diagnostics instead of fixed-input teacher forcing") parser.add_argument("--split", choices=("calibration", "validation", "heldout"), help="Select one split; screen candidates on calibration only") parser.add_argument("--max-records", type=int, help="Evaluate only the first N records after split filtering; never truncates tokens") parser.add_argument("--describe", action="store_true") parser.add_argument("--device-ownership-confirmed", action="store_true", help="Required for device execution; operator must first ensure no other process owns the P150") args = parser.parse_args() for name in ("weights", "input", "output", "cache_root", "runtime"): setattr(args, name, getattr(args, name).resolve()) if not 2 <= args.context_limit <= 8192 or args.cpu_threads < 1: raise ValueError("context-limit must be 2..8192 and cpu-threads positive") if not 0 <= args.generate_tokens <= 256: raise ValueError("generate-tokens must be in [0,256]") # Restrict every persistent artifact/cache write to this isolated workspace. for path in (args.output, args.cache_root): if not path.is_relative_to(ROOT) or path == ROOT or path.is_relative_to(args.weights): raise ValueError(f"Write path must be isolated under {ROOT}, never under weights: {path}") config_path = args.weights / "config.json" config = json.loads(config_path.read_text())["text_config"] if config["num_hidden_layers"] != 32 or config["hidden_size"] != 4096 or config["vocab_size"] != 248320: raise ValueError("This evaluator is scoped to the original Qwen3.5-9B architecture") records = load_records(args.input, config["vocab_size"], args.context_limit) if args.split: records = [record for record in records if record["split"] == args.split] if args.max_records is not None: if args.max_records < 1: raise ValueError("max-records must be positive") records = records[:args.max_records] if not records: raise ValueError("No records remain after split/record selection") if max(len(record["token_ids"]) for record in records) + args.generate_tokens > args.context_limit: raise ValueError("Prompt plus generation exceeds the declared context limit") inventory = json.loads((ROOT / "precision-plan.json").read_text()) overrides = json.loads(args.overrides.read_text()) if args.overrides else None precision = precision_map(args.profile, overrides, config["layer_types"]) environment = runtime_environment(args.profile, inventory, args.mtp) source_hashes = {str(path.relative_to(args.runtime)): digest(path) for path in sorted(args.runtime.rglob("*")) if path.is_file() and path.suffix in {".py", ".cpp", ".hpp", ".h"}} identity = {"precision": precision, "environment": environment, "sources": source_hashes, "config_sha256": digest(config_path), "declared_upstream_revision": REVISION} native_manifest_path = args.weights / "native_manifest.json" if native_manifest_path.is_file(): native_manifest = json.loads(native_manifest_path.read_text()) if native_manifest["precision"] != precision: raise ValueError("Requested precision differs from the native checkpoint") identity["native_manifest_sha256"] = digest(native_manifest_path) else: identity["index_sha256"] = digest(args.weights / "model.safetensors.index.json") candidate_hash = hashlib.sha256(canonical(identity).encode()).hexdigest() cache = args.cache_root / candidate_hash metadata = {"schema_version": 1, "backend": "ttnn-native", "status": "not_started", "profile": args.profile, "precision": precision, "candidate_sha256": candidate_hash, "input_sha256": digest(args.input), "vocab_size": config["vocab_size"], "selection": {"split": args.split, "max_records": args.max_records, "record_ids": [record["id"] for record in records]}, "max_seq_len": max(128, ((max(len(r["token_ids"]) for r in records) + args.generate_tokens + 127) // 128) * 128), "source_identity": identity, "cache_path": str(cache), "runtime_environment": environment, "alignment": "row p consumes token_ids[:p+1], predicts token_ids[p+1]; positions 0..N-2", "full_vocabulary": True, "mtp": args.mtp, "vision": False, "trace": False, "speculative_generation": False, "generation_max_new_tokens": args.generate_tokens, "likelihood_scope": "prompt targets plus self-selected greedy targets; not heldout NLL" if args.generate_tokens else "fixed teacher-forced input targets", "state_dtype": "bf16", "activation_dtype": "bf16", "kv_dtype": "bf16", "weight_loader": "native quantized checkpoint, evaluation-only unverified load" if native_manifest_path.is_file() else "read-only original text safetensors, preserving FP32 exceptions and optional lossless MTP", "evaluator_sha256": digest(Path(__file__)), "runtime_comparison_caveat": "BF16 weights retain BF16 arithmetic and native GDN recurrence/frontend. Generic controls share nonpacked matmul settings; packed/output/norm fusions are a separate runtime comparison.", "created_at": datetime.datetime.now(datetime.timezone.utc).isoformat()} if args.describe: print(json.dumps(metadata, indent=2, sort_keys=True)) return if not args.device_ownership_confirmed: raise RuntimeError("Refusing device access: use --describe now. Execution requires exclusive P150 ownership, not merely a second container.") if args.output.exists() and any(args.output.iterdir()): raise ValueError("Output directory must be new/empty; refusing to overwrite prior evaluation") args.output.mkdir(parents=True, exist_ok=True) cache.mkdir(parents=True, exist_ok=True) metadata["status"] = "running" save_json(args.output / "metadata.json", metadata) try: evaluate(args, records, metadata, cache, environment) except BaseException as error: metadata["status"] = "failed" metadata["error"] = f"{type(error).__name__}: {error}" raise finally: save_json(args.output / "metadata.json", metadata) if __name__ == "__main__": main()