#!/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()