test1111111 / tests /test_reference.py
spitfire4794's picture
deploy c66f5aa: decode push
421b8c2
Raw History Blame Contribute Delete
7.47 kB
"""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"]