#!/usr/bin/env python3 """Validate packed channel orders on complete LFM2 convolution paths.""" from __future__ import annotations import argparse import json from pathlib import Path import mlx.core as mx import numpy as np from probe_gated_path import cosine, load_tensor, qdq, relative_mse, short_conv_path def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--model", type=Path, required=True) parser.add_argument("--packing-dir", type=Path, required=True) parser.add_argument("--bits", type=int, default=4) parser.add_argument("--group-size", type=int, default=64) parser.add_argument("--samples", type=int, default=2) parser.add_argument("--sequence-length", type=int, default=32) parser.add_argument("--seed", type=int, default=20260811) parser.add_argument("--output", type=Path, required=True) args = parser.parse_args() packing_files = sorted( args.packing_dir.glob("packing-layer*-exact5k.json"), key=lambda path: int(path.stem.split("layer")[1].split("-")[0]), ) results = [] for packing_file in packing_files: packing = json.loads(packing_file.read_text()) layer = int(packing["layer"]) order = mx.array(packing["best"]["order"]) prefix = f"model.layers.{layer}.conv" in_weight = mx.array(load_tensor(args.model, f"{prefix}.in_proj.weight")) kernel = mx.array(load_tensor(args.model, f"{prefix}.conv.weight")[:, 0, :]) out_weight = mx.array(load_tensor(args.model, f"{prefix}.out_proj.weight")) hidden = out_weight.shape[0] mx.random.seed(args.seed + layer) x = mx.random.normal( (args.samples, args.sequence_length, hidden), dtype=mx.float32 ) x /= mx.sqrt(mx.mean(mx.square(x), axis=-1, keepdims=True) + 1e-6) reference = short_conv_path(x, in_weight, kernel, out_weight) baseline = short_conv_path( x, qdq(in_weight, args.bits, args.group_size), kernel, qdq(out_weight, args.bits, args.group_size), ) packed_in = mx.concatenate( [ in_weight[:hidden][order], in_weight[hidden : 2 * hidden][order], in_weight[2 * hidden :][order], ], axis=0, ) packed_kernel = kernel[order] packed_out = out_weight[:, order] exact = short_conv_path(x, packed_in, packed_kernel, packed_out) candidate = short_conv_path( x, qdq(packed_in, args.bits, args.group_size), packed_kernel, qdq(packed_out, args.bits, args.group_size), ) mx.eval(reference, baseline, exact, candidate) baseline_mse = relative_mse(reference, baseline) candidate_mse = relative_mse(reference, candidate) result = { "layer": layer, "packing_label": packing["best"]["label"], "fp_invariance_relative_mse": relative_mse(reference, exact), "baseline_relative_mse": baseline_mse, "packed_relative_mse": candidate_mse, "baseline_cosine": cosine(reference, baseline), "packed_cosine": cosine(reference, candidate), "path_mse_reduction_percent": 100.0 * (baseline_mse - candidate_mse) / baseline_mse, "weight_mse_reduction_percent": packing["mse_reduction_percent"], } results.append(result) print( f"layer={layer:02d} path_reduction={result['path_mse_reduction_percent']:+.3f}% " f"weight_reduction={result['weight_mse_reduction_percent']:+.3f}%" ) path_reductions = [item["path_mse_reduction_percent"] for item in results] weight_reductions = [item["weight_mse_reduction_percent"] for item in results] invariance = [item["fp_invariance_relative_mse"] for item in results] summary = { "layers": len(results), "path_improved_layers": sum(value > 0 for value in path_reductions), "mean_path_mse_reduction_percent": float(np.mean(path_reductions)), "median_path_mse_reduction_percent": float(np.median(path_reductions)), "mean_weight_mse_reduction_percent": float(np.mean(weight_reductions)), "max_fp_invariance_relative_mse": float(np.max(invariance)), } payload = {"summary": summary, "layers": results} args.output.parent.mkdir(parents=True, exist_ok=True) args.output.write_text(json.dumps(payload, indent=2) + "\n") print(json.dumps(summary, indent=2)) if __name__ == "__main__": main()