"""Offline real-Transformers oracle, independent of the existing NumPy oracle.""" import importlib.util import json from pathlib import Path import socket import numpy as np import pytest torch = pytest.importorskip("torch") from tokenizers import Tokenizer, models, pre_tokenizers from transformers import AutoModelForCausalLM, LlamaConfig, PreTrainedTokenizerFast, Qwen3Config from cism import Engine, _native from cism.loader import import_model spec = importlib.util.spec_from_file_location( "validate_reference", Path(__file__).resolve().parents[1] / "scripts" / "validate_reference.py", ) validation = importlib.util.module_from_spec(spec) spec.loader.exec_module(validation) @pytest.fixture(autouse=True) def offline(monkeypatch): def no_network(*args, **kwargs): raise AssertionError("Reference tests must not access the network") monkeypatch.setenv("HF_HUB_OFFLINE", "1") monkeypatch.setattr(socket.socket, "connect", no_network) previous_threads = torch.get_num_threads() torch.set_num_threads(1) yield torch.set_num_threads(previous_threads) @pytest.fixture(params=[("llama", False), ("llama", True), ("qwen3", False), ("qwen3", True)]) def checkpoint(tmp_path, request): architecture, tied = request.param config_type = LlamaConfig if architecture == "llama" else Qwen3Config config = config_type( vocab_size=43, hidden_size=40, intermediate_size=53, num_hidden_layers=2, num_attention_heads=4, num_key_value_heads=2, head_dim=12, max_position_embeddings=80, rms_norm_eps=1e-5, rope_theta=777777.0, tie_word_embeddings=tied, bos_token_id=1, eos_token_id=2, pad_token_id=0, attention_dropout=0.0, ) with torch.random.fork_rng(devices=[]): torch.manual_seed(312) model = AutoModelForCausalLM.from_config(config, attn_implementation="eager").float().eval() with torch.no_grad(): for parameter in model.parameters(): if parameter.ndim == 1: parameter.uniform_(0.75, 1.25) else: parameter.normal_(0, 0.12) model.save_pretrained(tmp_path) tokenizer = Tokenizer(models.WordLevel({f"t{i}": i for i in range(43)}, unk_token="t0")) tokenizer.pre_tokenizer = pre_tokenizers.Whitespace() PreTrainedTokenizerFast( tokenizer_object=tokenizer, unk_token="t0", bos_token="t1", eos_token="t2", pad_token="t0", ).save_pretrained(tmp_path) return tmp_path @pytest.mark.parametrize("precision", validation.PRECISIONS) def test_transformers_logits_and_greedy(checkpoint, precision): loaded = import_model(checkpoint, local_files_only=True) native = _native.Model(loaded.config, loaded.weights, precision) engine = Engine(native, loaded.tokenizer, loaded.config, str(checkpoint)) reference = validation.load_reference(loaded, precision) assert reference.config._attn_implementation == "eager" assert all(p.device.type == "cpu" and p.dtype == torch.float32 for p in reference.parameters()) if loaded.config["tie_word_embeddings"]: assert reference.get_input_embeddings().weight is reference.get_output_embeddings().weight for name, parameter in reference.named_parameters(): if parameter.ndim == 1: np.testing.assert_array_equal(parameter.detach().numpy(), loaded.weights[name]) prompt = [3, 7, 1, 19, 4, 8, 12, 6, 25, 9, 11, 5, 16, 31, 20, 10] for length in (4, 8, 16): # Hybrid 4-bit is a lossy format: gate logits at an absolute level # calibrated to catch structural bugs (measured: quant-rule noise # ~0.2 on adversarial tiny weights, row-stride-class bugs ~2.9). # Argmax agreement below still guards greedy quality per step. atol = 0.5 if precision == "hybrid-int4" else 8e-6 report = validation.compare(engine, native, reference, prompt[:length], max_gen=4, atol=atol) assert report["passed"], json.dumps(report, indent=2) assert len(report["native_greedy_token_ids"]) == 4 assert len(report["teacher_forced"]) == 4 np.testing.assert_array_equal(engine.logits(prompt[:length]), native.logits(prompt[:length])) @pytest.mark.parametrize("precision,qmax", [("int8", 127), ("hybrid-int4", 7)]) def test_reconstruction_rounding_and_row_blocks(precision, qmax): below = np.nextafter(np.float32(0.5), np.float32(0)) values = np.array([qmax, -qmax, 0.5, -0.5, 1.5, -1.5, below, -below], np.float32) original = np.zeros((3, 37), np.float32) original[0, :8] = values original[1, :8] = values * 2 original[0, 32:] = [qmax * 4, 2, -2, 6, -6] original[1, 32:] = [qmax * 8, 4, -4, 12, -12] snapshot = original.copy() actual = validation.reconstruct_weight(original, "model.layers.0.mlp.down_proj.weight", precision) if precision == "hybrid-int4": # Signed-max rule (matches the native packer): extreme maps exactly, # opposite tail clips to +-7/8 of the extreme (see the packer comment). expected = np.zeros_like(original) expected[0, :8] = [7, -6.125, 0.875, -0.875, 1.75, -1.75, 0.875, -0.875] expected[1, :8] = expected[0, :8] * 2 expected[0, 32:] = [28, 3.5, -3.5, 7, -7] expected[1, 32:] = expected[0, 32:] * 2 else: expected = np.zeros_like(original) expected[0, :8] = [128, -128, 0, 0, 0, 0, 0, 0] expected[1, :8] = expected[0, :8] * 2 expected[0, 32:] = [qmax * 4, 4, -4, 8, -8] expected[1, 32:] = expected[0, 32:] * 2 # Dedicated scale=1 half ties, including a float immediately below 0.5. np.testing.assert_array_equal( validation.reconstruct_weight(values[None], "lm_head.weight", precision), [[qmax, -qmax, 1, -1, 2, -2, 0, 0]], ) np.testing.assert_array_equal(actual, expected) np.testing.assert_array_equal(original, snapshot) np.testing.assert_array_equal(validation.reconstruct_weight(original, "x", "fp32"), original) np.testing.assert_array_equal(validation.reconstruct_weight(values, "model.norm.weight", precision), values) np.testing.assert_array_equal( validation.reconstruct_weight(original, "model.embed_tokens.weight", "hybrid-int4"), validation.reconstruct_weight(original, "model.embed_tokens.weight", "int8"), ) tiny = np.array([[0, np.finfo(np.float32).smallest_subnormal]], np.float32) np.testing.assert_array_equal(validation.reconstruct_weight(tiny, "lm_head.weight", precision), tiny) def test_cli_json_and_failure_status(checkpoint, capsys): args = [str(checkpoint), "--local-files-only", "--precision", "int8", "--token-lengths", "4", "8", "16", "--max-gen", "4", "--prompt", " ".join(f"t{i}" for i in range(3, 25))] assert validation.main(args) == 0 result = json.loads(capsys.readouterr().out) assert result["passed"] assert result["resolved_source"] == str(checkpoint.resolve()) assert result["reference"]["weights"] == "reconstructed-int8" assert result["reference"]["dtype"] == "float32" assert result["max_gen"] == 4 assert [case["token_length"] for case in result["cases"]] == [4, 8, 16] assert validation.main(args + ["--atol", "0"]) == 1 assert json.loads(capsys.readouterr().out)["passed"] is False assert validation.main(args + ["--max-gen", "-1"]) == 2 assert "positive" in json.loads(capsys.readouterr().out)["error"]