"""Whole original compiled decoder-run receipts, using the Rust smoke geometry. No downloads or real model weights. Run only under test-supervisor --command. The fixture uses real GDN/HC dimensions and two small MoE experts; it is a graph and numerical contract, not a model-throughput benchmark. """ import io import os from pathlib import Path import runpy import sys print("gdn-run reference: importing pinned runtime", file=sys.stderr, flush=True) import mlx.core as mx import mlx.nn as nn from mlx.utils import tree_flatten from mlx_lm.models.switch_layers import QuantizedSwitchLinear from mtplx.models import qwen4_exp as qwen def main(): assert mx.__version__ == "0.32.2", mx.__version__ helpers = runpy.run_path(str(Path(__file__).with_name("mtplx-kernel-fixtures.py"))) emit = helpers["emit"] # Explicitly mirror decoder.rs's current connected model fixture. os.environ.update(MTPLX_FUSED_HC="1", MTPLX_FUSED_HC_V3="0", MTPLX_FUSED_GDN_STEP="1", MTPLX_FUSED_GDN_CONVNORM="0", MTPLX_FUSED_CONVNORM_VERIFY="1", MTPLX_FUSED_GDN_OUT="0", MTPLX_FUSED_MOE_DECODE="0", MTPLX_FUSED_MOE_VERIFY="0") args = qwen.TextArgs(num_hidden_layers=1, vocab_size=32, num_experts=2, num_experts_per_tok=1, moe_intermediate_size=64, shared_expert_intermediate_size=64, layer_types=["linear_attention"]) model = qwen.Qwen4ExpTextModel(args) layer = model.layers[0] zero = lambda shape: mx.zeros(shape, mx.bfloat16) def quantized(n, k, experts=None): # Tiny constructors avoid evaluating random placeholder parameters. module = (nn.QuantizedLinear(64, 1, bias=False, group_size=64, bits=4) if experts is None else QuantizedSwitchLinear(64, 1, 1, bias=False, group_size=64, bits=4)) prefix = () if experts is None else (experts,) module.weight = mx.zeros((*prefix, n, k // 8), mx.uint32) module.scales = zero((*prefix, n, k // 64)) module.biases = zero((*prefix, n, k // 64)) return module norm = mx.ones((10240,), mx.bfloat16) down, up, inject = zero((320, 10240)), zero((10240, 320)), zero((4, 10240)) for hc in (layer.attn_hyper_connection, layer.mlp_hyper_connection): hc.hc_norm.weight = norm hc.input_mix_weight_down.weight = down hc.input_mix_weight_up.weight = up hc.block_inject_weight.weight = inject gdn = layer.linear_attn projected = quantized(16480, 2560) gdn.in_proj_fused = qwen._FusedGDNInProj(projected.weight, projected.scales, projected.biases, 64, 4, "affine", [10240, 16384, 16432]) for name in ("qkv", "z", "a", "b"): gdn.pop("in_proj_" + name, None) gdn.out_proj = quantized(2560, 6144) gdn.conv1d.weight = zero((10240, 4, 1)) gdn.A_log = zero((48,)) gdn.dt_bias = zero((48,)) gdn.norm.weight = mx.ones((128,), mx.bfloat16) experts = quantized(128, 2560, 2) layer.mlp.switch_mlp = qwen._FusedGateUpSwitchGLU(quantized(2560, 64, 2), experts.weight, experts.scales, experts.biases, 64, 4, "affine") shared = quantized(128, 2560) layer.mlp.shared_expert = qwen._FusedGateUpMLP(quantized(2560, 64), shared.weight, shared.scales, shared.biases, 64, 4, "affine") layer.mlp.gate = quantized(2, 2560) layer.mlp.shared_expert_gate = quantized(1, 2560) model.eval() assert not gdn.training print("gdn-run reference: materializing fixture parameters", file=sys.stderr, flush=True) mx.eval(*[a for _, a in tree_flatten(layer.parameters())]) for step, (rows, capture) in enumerate(2 * ((4, False), (1, False), (1, True), (4, True))): hyper_fused = step < 4 if step % 4 == 0: states = [zero((1, 3, 10240)), mx.zeros((1, 48, 128, 128), mx.float32)] os.environ["MTPLX_FUSED_HC"] = str(int(hyper_fused)) model._decode_run_fns.clear() print(f"gdn-run reference: step={step} rows={rows} capture={capture}", file=sys.stderr, flush=True) h = mx.full((1, rows, 10240), 0.125, mx.bfloat16) mx.eval(h, *states) token = qwen._VERIFY_CAPTURE.set(capture) try: outputs = model._get_run_fn((0,), capture)(h, *states) finally: qwen._VERIFY_CAPTURE.reset(token) dot = io.StringIO() mx.export_to_dot(dot, hidden=h, conv=states[0], delta=states[1], **{f"out{i}": a for i, a in enumerate(outputs)}) emit(f"gdn_run_r{rows}_capture{int(capture)}_hc{int(hyper_fused)}", outputs, rows=rows, capture=capture, hyper_fused=hyper_fused, dot=dot.getvalue(), shapes=[a.shape for a in outputs], contract=helpers["compiled_dot_contract"](dot.getvalue(), [f"out{i}" for i in range(len(outputs))], {"hidden", "conv", "delta"})) states = list(outputs[1:3]) if __name__ == "__main__": main()