hv-two-factor-lm / example.py
zeechimp's picture
Create example.py
b517ef0 verified
Raw History Blame Contribute Delete
2.95 kB
#!/usr/bin/env python3
"""
example.py — Load a trained two-factor LM and score sequences.
NumPy only.
"""
import numpy as np
from hv_two_factor_lm import (
HvTwoFactorLM,
make_sequence,
theory_p_marginal_full,
theory_r,
theory_csi_struct,
)
def load_model(path='two_factor_lm.npz'):
data = np.load(path)
model = HvTwoFactorLM(D=int(data['D']), K=int(data['K']),
decay=float(data['decay']))
model.token_hv = data['token_hv']
model.W_lm = data['W_lm']
model.b_lm = data['b_lm']
model.w_alpha = data['w_alpha']
model.w_inject = data['w_inject']
model.W_comp = data['W_comp']
model.W_struct = data['W_struct']
return model
def predict_marginal(model, alpha, inject, seq_len=32,
n_samples=256, seed=0):
"""Return model's batch-mean output distribution at (alpha, inject)."""
rng = np.random.default_rng(seed)
tokens = np.stack([
make_sequence(seq_len, alpha, inject, rng)
for _ in range(n_samples)])
h = model.encode(tokens)
a = np.full(n_samples, alpha, dtype=np.float32)
i = np.full(n_samples, inject, dtype=np.float32)
scores = model.lm_scores(h, a, i)
P = model.softmax(scores)
return P.mean(axis=(0, 1))
def main():
model = load_model('two_factor_lm.npz')
print(f"Loaded two_factor_lm.npz "
f"(D={model.D}, K={model.K}, decay={model.decay})")
print()
print("Batch-mean output distribution vs theoretical target:")
print(f" {'alpha':>6s} {'inject':>7s} "
f"{'model':<42s} {'target':<42s} {'KL':>7s}")
print(" " + "-" * 108)
for alpha in [0.35, 0.50, 0.65]:
for inject in [0.0, 0.15, 0.30]:
p_model = predict_marginal(model, alpha, inject)
p_target = theory_p_marginal_full(alpha, inject)
kl = float((p_model * np.log(
(p_model + 1e-12) / (p_target + 1e-12))).sum())
m = '[' + ' '.join(f'{v:.3f}' for v in p_model) + ']'
t = '[' + ' '.join(f'{v:.3f}' for v in p_target) + ']'
print(f" {alpha:>6.2f} {inject:>7.2f} "
f"{m:<42s} {t:<42s} {kl:>7.4f}")
print()
print("Ridge head outputs on a sample sequence:")
seq = make_sequence(32, 0.5, 0.15, np.random.default_rng(0))
h = model.encode(seq[None, :])
X = np.concatenate([h.reshape(-1, model.D),
np.ones((h.shape[1], 1), dtype=np.float32)], axis=-1)
r_pred = (X @ model.W_comp).mean(axis=0)
s_pred = float((X @ model.W_struct).mean())
r_true = theory_r(0.5)
s_true = np.log1p(theory_csi_struct(0.5, 0.15))
print(f" comp pred = [{', '.join(f'{v:.3f}' for v in r_pred)}]")
print(f" comp true = [{', '.join(f'{v:.3f}' for v in r_true)}]")
print(f" struct pred = {s_pred:.4f} (log1p space)")
print(f" struct true = {s_true:.4f} (log1p space)")
if __name__ == '__main__':
main()