Warm-Started Exact Rank Pruning / experiment.py
Mechanism confirmed, baseline not beaten
1import json
2import time
3import numpy as np
4
5
6def joint_prune(z_u, z_v, eta, lam):
7 """Exact proximal map for a shared paired 2-lambda column penalty."""
8 keep = (np.sum(z_u * z_u, axis=0) + np.sum(z_v * z_v, axis=0)) > 4.0 * eta * lam
9 out_u, out_v = z_u.copy(), z_v.copy()
10 out_u[:, ~keep] = 0.0
11 out_v[:, ~keep] = 0.0
12 return out_u, out_v, keep
13
14
15def balance(u, v, eps=0.0):
16 """Reciprocal column rescaling; eps=0 is the exact mathematical map."""
17 nu = np.linalg.norm(u, axis=0)
18 nv = np.linalg.norm(v, axis=0)
19 a = np.sqrt(nv / (nu + eps))
20 return u * a[None, :], v / a[None, :]
21
22
23def factor_step(u, v, x, eta, lam, mu):
24 residual = u @ v.T - x
25 gu = residual @ v + mu * u
26 z_u = u - eta * gu
27 # alternating update as in the stated formula
28 gv = (z_u @ v.T - x).T @ z_u + mu * v
29 z_v = v - eta * gv
30 u, v, keep = joint_prune(z_u, z_v, eta, lam)
31 if np.any(keep):
32 # balancing only retained pairs avoids 0/0 and preserves the product
33 ub, vb = balance(u[:, keep], v[:, keep])
34 u[:, keep], v[:, keep] = ub, vb
35 return u, v, keep
36
37
38def run_path(x, u0, v0, lambdas, eta=0.01, mu=1e-3, steps=250):
39 u, v = u0.copy(), v0.copy()
40 rows = []
41 for lam in lambdas:
42 for _ in range(steps):
43 u, v, _ = factor_step(u, v, x, eta, lam, mu)
44 active = np.linalg.norm(u, axis=0) > 1e-10
45 rows.append({
46 "lambda": float(lam), "rank": int(active.sum()),
47 "loss": float(0.5 * np.sum((x - u @ v.T) ** 2)),
48 "objective": float(0.5 * np.sum((x - u @ v.T) ** 2) +
49 0.5 * mu * (np.sum(u*u) + np.sum(v*v)) + 2 * lam * active.sum())
50 })
51 return u, v, rows
52
53
54def main():
55 rng = np.random.default_rng(1352)
56 # Low-rank positive target with deliberately excessive factor rank.
57 n, m, true_rank, r = 16, 13, 3, 8
58 x = rng.normal(size=(n, true_rank)) @ rng.normal(size=(m, true_rank)).T
59 x /= np.linalg.norm(x, 'fro') / np.sqrt(n*m)
60 # SVD initialization gives useful columns plus weak redundant columns.
61 p, s, qt = np.linalg.svd(x, full_matrices=False)
62 u0 = np.zeros((n, r)); v0 = np.zeros((m, r))
63 for j in range(r):
64 if j < len(s):
65 u0[:, j] = p[:, j] * np.sqrt(max(s[j], 1e-12))
66 v0[:, j] = qt[j, :] * np.sqrt(max(s[j], 1e-12))
67 else:
68 u0[:, j] = 0.03 * rng.normal(size=n)
69 v0[:, j] = 0.03 * rng.normal(size=m)
70 # Make the surplus columns genuinely weak but nonzero.
71 u0[:, true_rank:] *= 0.12
72 v0[:, true_rank:] *= 0.12
73
74 # Prediction 1: exact joint threshold boundary, swept over scales/lambdas.
75 eta = 0.01
76 threshold_errors = []
77 boundary_cases = []
78 for lam in np.geomspace(1e-4, 2.0, 24):
79 threshold = 4 * eta * lam
80 for scale in np.geomspace(0.05, 20.0, 17):
81 zu = np.array([[scale, 1.0]])
82 zv = np.array([[1.0, scale]])
83 q = np.sum(zu*zu, axis=0) + np.sum(zv*zv, axis=0)
84 _, _, keep = joint_prune(zu, zv, eta, lam)
85 threshold_errors.append(abs(float(keep[0]) - float(q[0] > threshold)))
86 if abs(q[0] - threshold) < max(threshold * 0.02, 1e-12):
87 boundary_cases.append((q[0], threshold, bool(keep[0])))
88 prox_check = {"mismatches": int(sum(threshold_errors)),
89 "tested": len(threshold_errors),
90 "near_boundary_cases": boundary_cases[:4]}
91
92 # Prediction 2: reciprocal scaling must preserve product and equalize norms.
93 products, norm_gaps, scale_invariance = [], [], []
94 for _ in range(100):
95 uu = rng.normal(size=(7, 1)); vv = rng.normal(size=(6, 1))
96 c = 10 ** rng.uniform(-6, 6)
97 uu *= c; vv /= c
98 ub, vb = balance(uu, vv, eps=0.0)
99 products.append(np.linalg.norm(uu @ vv.T - ub @ vb.T, 'fro'))
100 norm_gaps.append(abs(np.linalg.norm(ub) - np.linalg.norm(vb)))
101 scale_invariance.append(abs(np.linalg.norm(ub @ vb.T, 'fro') - np.linalg.norm(uu @ vv.T, 'fro')))
102 balance_check = {"max_product_error": float(max(products)),
103 "max_balanced_norm_gap": float(max(norm_gaps)),
104 "max_product_norm_error": float(max(scale_invariance))}
105
106 # Prediction 3: increasing lambda on a warm start removes columns monotonically.
107 lambdas = [0.0, 0.02, 0.08, 0.2, 0.5, 1.0, 2.0, 4.0]
108 t0 = time.perf_counter()
109 _, _, path = run_path(x, u0, v0, lambdas, eta=eta, mu=1e-3, steps=300)
110 warm_time = time.perf_counter() - t0
111 ranks = [row["rank"] for row in path]
112 monotone = all(ranks[i+1] <= ranks[i] for i in range(len(ranks)-1))
113
114 # Standard comparison: independently optimized fixed-rank factorizations.
115 baseline = []
116 for rr in [true_rank, r]:
117 ub, vb = u0[:, :rr].copy(), v0[:, :rr].copy()
118 t1 = time.perf_counter()
119 for _ in range(sum([300] * len(lambdas))):
120 # lambda=0 is ordinary factorized least squares with scale regularization
121 ub, vb, _ = factor_step(ub, vb, x, eta, 0.0, 1e-3)
122 baseline.append({"rank": rr, "loss": float(0.5*np.sum((x-ub@vb.T)**2)),
123 "seconds": time.perf_counter()-t1})
124 selected = min((row for row in path if row["loss"] <= baseline[1]["loss"] * 1.05),
125 key=lambda row: row["rank"], default=path[-1])
126 result = {
127 "prox_prediction": prox_check,
128 "balance_prediction": balance_check,
129 "path_prediction": {"lambdas": lambdas, "ranks": ranks,
130 "monotone_nonincreasing": monotone,
131 "losses": [row["loss"] for row in path],
132 "warm_seconds": warm_time},
133 "baseline": baseline,
134 "prediction_summary": {
135 "threshold_rule": {"predicted_mismatches": 0, "observed_mismatches": int(sum(threshold_errors)), "tested": len(threshold_errors)},
136 "balancing_rule": {"predicted_product_error": 0.0, "observed_max_product_error": float(max(products)), "predicted_norm_gap": 0.0, "observed_max_norm_gap": float(max(norm_gaps))},
137 "lambda_zero_no_pruning": {"predicted_rank": r, "observed_rank": ranks[0]},
138 "warm_path_rank": {"predicted_nonincreasing": True, "observed_nonincreasing": monotone, "observed_transition": f"{ranks[0]}->{ranks[-1]}"}
139 },
140 "selected_path_point": selected,
141 "setup": {"shape": [n,m], "true_rank": true_rank, "max_rank": r,
142 "eta": eta, "steps_per_stage": 300}
143 }
144 print(json.dumps(result, indent=2))
145
146
147if __name__ == '__main__':
148 main()