Files
TOIFAS/code/diag_maskdegen.py
KiHoLee e243b29d29 Audit round: stored degeneracy measurement, figure guards, 44 assertions
diag_maskdegen writes data/maskdegen.csv so the claim it supports is
traceable; the legend guard inflates by the marker radius and refuses a
legend that leaves the canvas; one legend size on every figure.
2026-08-19 15:08:38 +09:00

80 lines
2.8 KiB
Python

# -*- coding: utf-8 -*-
"""Does unconstrained key training still degenerate at the main configuration?
The manuscript justifies fixing the keys by a measured failure: with the
keys free, training drives them to disjoint sparse supports, which is an
orthogonal slot allocation rather than a superposition, and which shrinks
the key space to the choice of a support. That was measured at d=64 and
has to be re-measured whenever the configuration moves, because it is
the reason the structured family is the main one.
Reported per user key: the number of entries holding 99 percent of the
energy, and the pairwise overlap of those supports. A dense key spreads
its energy over most of the L entries and the supports coincide; a
degenerate one concentrates on a few and the supports are disjoint.
"""
import sys
from pathlib import Path
import torch
sys.path.insert(0, str(Path(__file__).resolve().parent))
from exp_full import get_model, main_model, MAIN_D
def support99(w):
"""Smallest set of entries carrying 99 percent of the key energy."""
e = w.pow(2)
order = torch.argsort(e, descending=True)
c = torch.cumsum(e[order], 0) / e.sum()
k = int((c < 0.99).sum()) + 1
return set(order[:k].tolist()), k
def describe(name, W, rows=None):
L = W.shape[1]
sups, ks = [], []
for u in range(W.shape[0]):
sup, k = support99(W[u])
sups.append(sup)
ks.append(k)
ov = []
for i in range(len(sups)):
for j in range(i + 1, len(sups)):
ov.append(len(sups[i] & sups[j]) / max(1, min(len(sups[i]),
len(sups[j]))))
mo = sum(ov) / len(ov)
print("%-14s L=%3d 99%%-energy entries per key: %s "
"mean pairwise support overlap %.2f" % (name, L, ks, mo))
if rows is not None:
rows.append([name, L, "/".join(str(k) for k in ks), "%.4f" % mo])
def write_rows(rows):
"""Store the measurement so the manuscript sentence it justifies is
traceable to an artifact in data/ like every other quoted number."""
import csv
out = Path(__file__).resolve().parents[1] / "data" / "maskdegen.csv"
with open(out, "w", newline="") as f:
w = csv.writer(f)
w.writerow(["family", "L", "support99_per_key", "mean_overlap"])
w.writerows(rows)
print("[csv]", out)
def main():
print("main configuration d=%d" % MAIN_D)
rows = []
m_free = get_model(iters=4000) # keys learned, nothing frozen
describe("learned", m_free.masks().detach().cpu(), rows)
m_fix = main_model()
describe("Walsh-Hadamard", m_fix.masks().detach().cpu(), rows)
write_rows(rows)
print()
print("A degenerate key set shows few entries per key and near-zero")
print("overlap; a dense one shows most entries and overlap near one.")
if __name__ == "__main__":
main()