File size: 13,544 Bytes
28a1a01
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
421b8c2
 
 
 
 
 
 
 
 
 
 
28a1a01
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
"""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())