Warm-Started Exact Rank Pruning / experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()