File size: 4,635 Bytes
48883b3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/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()