"""Generate model-free GDN boundary fixtures with the installed MTPLX/MLX. Run with MTPLX's Python environment; this never loads or downloads a model. Large-row math follows GatedDeltaNet.__call__ in the pinned qwen4_exp.py; small rows call MTPLX's actual fused kernel. All files are diagnostic outputs. """ import argparse import hashlib import json import statistics import time from pathlib import Path import mlx.core as mx import mlx.nn as nn import numpy as np from mtplx.kernels.gdn_conv_norm import fused_gdn_conv_norm_rows from mtplx.models import qwen4_exp def delta_reference(output): from mlx_lm.models import gated_delta rng = np.random.default_rng(12345) initial = mx.array(rng.normal(0, 0.01, (1, 48, 128, 128)).astype(np.float32)) np.array(initial).tofile(output / "delta-initial.f32") source = Path(gated_delta.__file__).resolve() receipt = {"mlx_version": mx.__version__, "reference_source": str(source), "reference_sha256": hashlib.sha256(source.read_bytes()).hexdigest(), "cases": []} for rows in (1, 2, 3, 4, 5, 6, 7, 32, 2048): def normalized(scale): x = rng.normal(size=(1, rows, 16, 128)).astype(np.float32) return mx.array(x / np.linalg.norm(x, axis=-1, keepdims=True) * scale).astype(mx.bfloat16) q, k = normalized(128 ** -0.5), normalized(1) v = mx.array(rng.normal(0, 0.2, (1, rows, 48, 128)).astype(np.float32)).astype(mx.bfloat16) g = mx.array(rng.uniform(0.8, 0.999, (1, rows, 48)).astype(np.float32)) beta = mx.array(rng.uniform(0.1, 0.9, (1, rows, 48)).astype(np.float32)).astype(mx.bfloat16) mx.eval(q, k, v, g, beta, initial) def forward(): return gated_delta.gated_delta_kernel(q, k, v, g, beta, initial) mx.eval(*forward()) elapsed = [] for _ in range(10): start = time.perf_counter() result = forward() mx.eval(*result) elapsed.append((time.perf_counter() - start) * 1000) hashes = {} for name, array in zip(("q", "k", "v", "g", "beta", "out", "state"), (q, k, v, g, beta, *result)): raw = np.array(array.astype(mx.float32)).astype("