File size: 2,954 Bytes
b517ef0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
#!/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()