hv-locality / example.py
zeechimp's picture
Create example.py
27b3e6d verified
Raw History Blame Contribute Delete
3.07 kB
"""
Minimal usage example for hv-locality.
python example.py
Diagnoses three functions at the same I/O dimensions, prints a verdict
matrix, and profiles two synthetic tasks.
"""
import numpy as np
from hv_locality import (
HVLocality,
task_profile,
_identity_encoder_factory,
_weak_avalanche_factory,
_random_hash_factory,
)
N_IN = 64
N_OUT = 512
def _verdict_row(m, reports, funcs, k_intra):
row = [f"{k_intra:>8}"]
for name, _ in funcs:
v = m.verdict(reports[name], k_intra)
row.append(f"{v['tier']:>16}")
return " ".join(row)
def main() -> None:
m = HVLocality()
funcs = [
("identity", _identity_encoder_factory(N_OUT)),
("weak_avalanche", _weak_avalanche_factory(N_OUT, mix_rate=0.1)),
("random_hash", _random_hash_factory(N_OUT)),
]
reports = {}
print("=" * 70)
print("Diagnostic summary")
print("=" * 70)
print(f" {'function':<18} {'ε':>7} {'L(1)':>7} "
f"{'L(10)':>7} {'k*':>4} {'L_∞':>7}")
print(" " + "-" * 62)
for name, f in funcs:
rpt = m.diagnose(f, N_IN, N_OUT, function_name=name, seed=0)
reports[name] = rpt
print(f" {name:<18} {rpt.avalanche_epsilon:>7.4f} "
f"{rpt.L_at_1:>7.4f} {rpt.L_at_10:>7.4f} "
f"{rpt.saturation_k:>4} {rpt.plateau_value:>7.4f}")
print()
print("=" * 70)
print("Verdict matrix")
print("=" * 70)
print(f" {'k_intra':>8} " + " ".join(f"{n:>16}" for n, _ in funcs))
print(" " + "-" * (8 + 18 * len(funcs)))
for k in [1, 5, 10, 20, 50]:
print(_verdict_row(m, reports, funcs, k))
print()
print("=" * 70)
print("Task profile: synthetic 10-class 64-dim")
print("=" * 70)
rng = np.random.default_rng(42)
X_list, y_list = [], []
for c in range(10):
center = rng.standard_normal(64) * 2.0
X_c = center + rng.standard_normal((100, 64)) * 1.0
X_list.append(X_c)
y_list.append(np.full(100, c))
X = np.concatenate(X_list, axis=0).astype(np.float32)
y = np.concatenate(y_list, axis=0)
prof = task_profile(X, y, n_bits=64, seed=0)
print(f" k_intra (median) : {prof['k_intra']['median']:.2f}")
print(f" k_inter (median) : {prof['k_inter']['median']:.2f}")
print()
print("=" * 70)
print("Task profile: MNIST-like 8x8 bit patterns")
print("=" * 70)
rng = np.random.default_rng(7)
X_list, y_list = [], []
for c in range(10):
proto = rng.integers(0, 2, size=64).astype(np.float32)
noise = (rng.random((50, 64)) < 0.15).astype(np.float32)
X_c = np.clip(proto + noise - 2 * proto * noise, 0, 1)
X_list.append(X_c)
y_list.append(np.full(50, c))
X = np.concatenate(X_list, axis=0)
y = np.concatenate(y_list, axis=0)
prof = task_profile(X, y, n_bits=64, seed=0)
print(f" k_intra (median) : {prof['k_intra']['median']:.2f}")
print(f" k_inter (median) : {prof['k_inter']['median']:.2f}")
if __name__ == "__main__":
main()