File size: 3,112 Bytes
3fffa60
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
subspace_angle.py -- principal angles between BASE and OURS truncation subspaces.

Efficiency / subspace-drift evidence: BASE and OURS run the identical whitening
math at the same cost; the ONLY thing that changes is which subspace the
calibration selects. This quantifies that difference per layer.

For each matching layer we take the retained input subspace = row space of the
B factor (k x in), orthonormalize both, and compute principal angles via the SVD
of Q_base^T Q_ours (its singular values are cos of the principal angles).
Reports mean/max angle (degrees) per layer + an overall summary to
results/efficiency/.
"""
import os
import sys
import glob
import argparse
import math
import torch

sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import common as C


def ortho_rowspace(B):
    """Return an orthonormal basis (in x k) of the row space of B (k x in)."""
    # columns of Q span row space of B
    Q, _ = torch.linalg.qr(B.t().to(torch.float64), mode="reduced")
    return Q


def principal_angles_deg(B1, B2):
    Q1 = ortho_rowspace(B1)
    Q2 = ortho_rowspace(B2)
    M = Q1.t() @ Q2
    s = torch.linalg.svdvals(M).clamp(-1.0, 1.0)
    angles = torch.arccos(s) * (180.0 / math.pi)
    return angles


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--base_weights", required=True)
    ap.add_argument("--ours_weights", required=True)
    ap.add_argument("--out", default=None)
    args = ap.parse_args()

    base_Bs = sorted(glob.glob(os.path.join(args.base_weights, "*_B.pt")))
    per_layer = {}
    all_means = []
    for bp in base_Bs:
        fn = os.path.basename(bp)
        op = os.path.join(args.ours_weights, fn)
        if not os.path.exists(op):
            continue
        B1 = torch.load(bp, map_location="cpu").float()
        B2 = torch.load(op, map_location="cpu").float()
        if B1.shape != B2.shape:
            per_layer[fn] = {"skipped": f"shape {tuple(B1.shape)} vs {tuple(B2.shape)}"}
            continue
        ang = principal_angles_deg(B1, B2)
        per_layer[fn] = {"k": B1.shape[0], "mean_deg": float(ang.mean()),
                         "max_deg": float(ang.max()), "median_deg": float(ang.median())}
        all_means.append(float(ang.mean()))

    summary = {
        "base_weights": args.base_weights,
        "ours_weights": args.ours_weights,
        "n_layers": len(all_means),
        "overall_mean_deg": sum(all_means) / len(all_means) if all_means else None,
        "overall_min_layer_mean_deg": min(all_means) if all_means else None,
        "overall_max_layer_mean_deg": max(all_means) if all_means else None,
        "git_hash": C.git_hash(),
        "per_layer": per_layer,
    }
    root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
    out = args.out or os.path.join(root, "results", "efficiency",
                                   f"{C.git_hash()}_subspace_angle.json")
    C.dump_json(summary, out)
    print(f"[subspace] {len(all_means)} layers  "
          f"overall_mean={summary['overall_mean_deg']}deg -> {out}")


if __name__ == "__main__":
    main()