"""Compare installed CISM with real Transformers eager, FP32 CPU inference. Run with the environment containing the installed native extension, for example: .venv/Scripts/python scripts/validate_reference.py SupraLabs/Supra-Mini-v5-8M \ --precision int8 --local-files-only Quantized parity uses reconstructed weights, NOT the original float checkpoint. HF still computes in FP32; native quantized accumulation has a different order, so bitwise equality is not expected. Exit codes: 0 parity, 1 mismatch, 2 error. """ from __future__ import annotations import argparse from contextlib import redirect_stdout from importlib.metadata import version import json import math import sys import numpy as np PRECISIONS = ("fp32", "int8", "hybrid-int4", "hybrid-fp4") DEFAULT_PROMPT = ( "The future of artificial intelligence is an interesting subject. " "Scientists study how computers learn and how people can use them to solve problems." ) E2M1 = np.array([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=np.float64) def encode_e4m3(values): """Round-to-nearest-even positive floats to E4M3, mirroring the native encoder.""" values = np.asarray(values, dtype=np.float32) bits = np.zeros(values.shape, dtype=np.int64) normal = values >= np.float32(2.0 ** -6) with np.errstate(divide="ignore"): exponent = np.floor(np.log2(np.where(normal, values, np.float32(1)))).astype(np.int64) mantissa = values / (2.0 ** exponent) quantized = np.rint((mantissa - np.float32(1)) * 8).astype(np.int64) field = exponent + 7 + (quantized >= 8) quantized = np.where(quantized >= 8, 0, quantized) bits[normal] = field[normal] * 8 + quantized[normal] subnormal = np.rint(values * 512.0).astype(np.int64) bits[~normal] = np.minimum(subnormal[~normal], 0x7E) bits = np.minimum(bits, 0x7E) decoded = np.where(bits >= 8, (1 + (bits & 7) / 8) * 2.0 ** ((bits >> 3) - 7), (bits & 15) * 2.0 ** -9) return decoded.astype(np.float32) def reconstruct_fp4(weight): """Decode E2M1 + E4M3-per-16-elements packing (runtime.cpp Storage::fp4).""" original = np.asarray(weight, dtype=np.float32) rows, cols = original.shape pad = (-cols) % 16 blocks = np.pad(original, ((0, 0), (0, pad))).reshape(rows, -1, 16) amax = np.abs(blocks).max(axis=-1) ideal = np.where(amax == 0, np.float32(0), amax / 6) scale = encode_e4m3(ideal) magnitude = np.abs(blocks) / scale[..., None] index = np.argmin(np.abs(magnitude[..., None] - E2M1), axis=-1) decoded = np.where(np.sign(blocks) < 0, -E2M1[index], E2M1[index]) * scale[..., None] return decoded.reshape(rows, -1)[:, :cols].astype(np.float32) def reconstruct_weight(weight, name, precision): """Decode runtime.cpp Matrix packing, with FP32 scales and per-row blocks.""" if precision not in PRECISIONS: raise ValueError(f"Unknown precision: {precision}") original = np.asarray(weight, dtype=np.float32) if precision == "fp32" or original.ndim == 1: return original.copy() if precision == "hybrid-fp4" and ".mlp." in name: return reconstruct_fp4(original) int4 = precision == "hybrid-int4" and ".mlp." in name block_size = 32 if int4 else original.shape[1] qmax = np.float32(7 if int4 else 127) decoded = np.empty_like(original) for start in range(0, original.shape[1], block_size): values = original[:, start : start + block_size] if int4: # Signed-max (Q4_0 rule, matches the native packer): the extreme # keeps its sign (first max wins, like the native strict-greater # scan), d = extreme/-8, truncating quant, codes map to (code-8). index = np.argmax(np.abs(values), axis=1) extreme = values[np.arange(values.shape[0]), index][:, None] d = np.where(extreme == 0, np.float32(1.0), extreme / np.float32(-8.0)) xi = (values / d + np.float32(8.5)).astype(np.int32) xi = np.clip(xi, 0, 15) decoded[:, start : start + block_size] = (xi.astype(np.float32) - np.float32(8.0)) * d continue maximum = np.max(np.abs(values), axis=1, keepdims=True) scale = np.where(maximum == 0, np.float32(1), maximum / qmax) scale = np.maximum(scale, np.finfo(np.float32).smallest_subnormal) divided = values / scale # np.round uses ties-to-even. Adding 0.5 in FP32 also misrounds values # immediately below a half; splitting the fraction matches std::round. fraction, integral = np.modf(np.abs(divided)) rounded = np.copysign(integral + (fraction >= np.float32(0.5)), divided) decoded[:, start : start + block_size] = np.clip(rounded, -qmax, qmax) * scale return decoded def load_reference(loaded, precision): import torch from transformers import AutoModelForCausalLM reference = AutoModelForCausalLM.from_pretrained( loaded.source, local_files_only=True, trust_remote_code=False, dtype=torch.float32, attn_implementation="eager", ).to("cpu").eval() if reference.config.model_type != loaded.config["model_type"]: raise ValueError("CISM and Transformers resolved different architectures") state = reference.state_dict() if set(state) != set(loaded.weights): raise ValueError("CISM and Transformers loaded different weight names") for name, parameter in state.items(): if not np.array_equal(parameter.detach().numpy(), loaded.weights[name]): raise ValueError(f"CISM and Transformers loaded different source weights: {name}") if loaded.config["tie_word_embeddings"]: if reference.get_input_embeddings().weight is not reference.get_output_embeddings().weight: raise ValueError("Transformers did not preserve tied embeddings/head") if precision != "fp32": with torch.no_grad(): # named_parameters deduplicates tied weights. Reconstruct from HF's # own loaded values, not CISM's arrays, and never quantize twice. for name, parameter in reference.named_parameters(): parameter.copy_(torch.from_numpy(reconstruct_weight( parameter.detach().numpy(), name, precision, ))) return reference def compare(engine, native, reference, prompt, max_gen=4, atol=1e-4): """Teacher-force HF's continuation; separately exercise native cached decode.""" import torch expected = [] measurements = [] with torch.inference_mode(): for step in range(max_gen): prefix = prompt + expected target = reference( input_ids=torch.tensor([prefix], dtype=torch.long, device="cpu"), use_cache=False, ).logits[0, -1].float().numpy() actual = engine.logits(prefix) if actual.shape != target.shape or not (np.isfinite(actual).all() and np.isfinite(target).all()): raise ValueError("Invalid shape or nonfinite last-token logits") error = np.abs(actual.astype(np.float64) - target.astype(np.float64)) native_id, reference_id = int(actual.argmax()), int(target.argmax()) measurements.append({ "step": step, "prefix_length": len(prefix), "max_abs_error": float(error.max()), "mean_abs_error": float(error.mean()), "native_argmax_id": native_id, "reference_argmax_id": reference_id, "within_tolerance": bool(error.max() <= atol), }) expected.append(reference_id) # Ignore EOS on both sides for exactly max_gen raw greedy tokens. No HF # generate() processors, repetition penalties, or generation_config defaults. session = native.create_session(prompt, max_new_tokens=max_gen, temperature=0.0, eos_token_ids=[]) actual_tokens = session.next_tokens(1) + session.next_tokens(max_gen - 1) eos = engine.config.get("eos_token_id", getattr(engine.tokenizer, "eos_token_id", None)) eos_ids = [] if eos is None else eos if isinstance(eos, list) else [eos] engine_expected = [] for token in expected: engine_expected.append(token) if token in eos_ids: break engine_result = engine.generate(prompt, max_new_tokens=max_gen, temperature=0.0) greedy_match = actual_tokens == expected engine_match = engine_result.token_ids == engine_expected return { "prompt_token_ids": prompt, "token_length": len(prompt), "teacher_forced": measurements, "max_abs_error": max(item["max_abs_error"] for item in measurements), "mean_abs_error": float(np.mean([item["mean_abs_error"] for item in measurements])), "native_greedy_token_ids": actual_tokens, "reference_greedy_token_ids": expected, "greedy_match": greedy_match, "engine_token_ids": engine_result.token_ids, "reference_eos_stopped_token_ids": engine_expected, "engine_finish_reason": engine_result.finish_reason, "engine_match": engine_match, "passed": greedy_match and engine_match and all( item["within_tolerance"] and item["native_argmax_id"] == item["reference_argmax_id"] for item in measurements ), } def validate_model(model, *, revision=None, precision="fp32", token_lengths=(4, 8, 16), max_gen=4, prompt=DEFAULT_PROMPT, local_files_only=False, atol=1e-4): import torch import transformers import cism from cism import Engine, _native from cism.loader import import_model if precision not in PRECISIONS: raise ValueError(f"Unknown precision: {precision}") if not token_lengths or any(length < 1 for length in token_lengths) or max_gen < 1: raise ValueError("Token lengths and max_gen must be positive") if not math.isfinite(atol) or atol < 0: raise ValueError("atol must be finite and nonnegative") loaded = import_model(model, revision=revision, local_files_only=local_files_only) tokens = list(loaded.tokenizer.encode(prompt, add_special_tokens=True)) if len(tokens) < max(token_lengths): raise ValueError(f"Prompt has only {len(tokens)} tokens; supply a longer --prompt") if max(token_lengths) + max_gen > loaded.config["max_position_embeddings"]: raise ValueError("Token length + max_gen exceeds model context") native = _native.Model(loaded.config, loaded.weights, precision) engine = Engine(native, loaded.tokenizer, loaded.config, str(model), revision=loaded.revision) reference = load_reference(loaded, precision) cases = [compare(engine, native, reference, tokens[:length], max_gen, atol) for length in token_lengths] return { "model": str(model), "requested_revision": revision, "resolved_source": loaded.source, "resolved_revision": loaded.revision, "architecture": loaded.config["model_type"], "reference_class": type(reference).__name__, "precision": precision, "native_info": native.info, "reference": { "device": "cpu", "dtype": "float32", "attention": reference.config._attn_implementation, "use_cache": False, "weights": "original-fp32" if precision == "fp32" else f"reconstructed-{precision}", "quantization": native.info["quantization"], "rounding": "none" if precision == "fp32" else "std::round (ties away from zero)", "norms": "unchanged-fp32", "tied_embeddings": loaded.config["tie_word_embeddings"], }, "versions": {"cism": version("cism"), "torch": torch.__version__, "transformers": transformers.__version__}, "installed_paths": {"cism": cism.__file__, "native": _native.__file__}, "atol": atol, "max_gen": max_gen, "generation_eos_policy": "raw native/HF ignore EOS; Engine honors EOS", "max_abs_error": max(case["max_abs_error"] for case in cases), "mean_abs_error": float(np.mean([case["mean_abs_error"] for case in cases])), "passed": all(case["passed"] for case in cases), "cases": cases, } def main(argv=None): parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) parser.add_argument("model", help="HF model ID or local safetensors directory") parser.add_argument("--revision", help="HF revision resolved once by CISM's loader") parser.add_argument("--precision", choices=PRECISIONS, default="fp32") parser.add_argument("--token-lengths", "--token-length", type=int, nargs="+", default=[4, 8, 16]) parser.add_argument("--max-gen", type=int, default=4) parser.add_argument("--prompt", default=DEFAULT_PROMPT) parser.add_argument("--local-files-only", action="store_true") parser.add_argument("--atol", type=float, default=1e-4, help="Maximum absolute logit error (default: 1e-4)") args = parser.parse_args(argv) try: # Keep stdout machine-readable even if dependency loading prints messages. with redirect_stdout(sys.stderr): import torch torch.set_num_threads(1) result = validate_model(**vars(args)) status = 0 if result["passed"] else 1 except Exception as exc: result = {"model": args.model, "precision": args.precision, "passed": False, "error": f"{type(exc).__name__}: {exc}"} status = 2 print(json.dumps(result, indent=2, allow_nan=False)) return status if __name__ == "__main__": raise SystemExit(main())