torch.manual_seed(0)
d_mem, h_mem, num_pairs = 32, 128, 16
k_pairs = nn.functional.normalize(torch.randn(num_pairs, d_mem), dim=-1)
v_pairs = nn.functional.normalize(torch.randn(num_pairs, d_mem), dim=-1)
order = torch.randperm(num_pairs * 4) % num_pairs # Each pair recurs 4x
memory = nn.Sequential(nn.Linear(d_mem, h_mem), nn.GELU(),
nn.Linear(h_mem, d_mem))
with torch.no_grad():
memory[0].weight.normal_(0, 1.0) # Unit-variance hidden units
memory[0].bias.zero_()
memory[2].weight.zero_() # Start empty: retrieve 0 everywhere
memory[2].bias.zero_()
def recall_loss(params, k, v):
return ((functional_call(memory, params, (k[None],))[0] - v)**2).sum()
grad_fn = grad(recall_loss)
def retrieval_mse(params):
with torch.no_grad():
pred = functional_call(memory, params, (k_pairs,))
return float(((pred - v_pairs)**2).mean())
params = {name: p.detach().clone() for name, p in memory.named_parameters()}
velocity = {name: torch.zeros_like(p) for name, p in params.items()}
before = retrieval_mse(params)
theta, eta = 0.005, 0.5
for t in order.tolist():
g = grad_fn(params, k_pairs[t], v_pairs[t]) # Surprise
for name in params:
velocity[name] = eta * velocity[name] - theta * g[name]
params[name] = params[name] + velocity[name]
after = retrieval_mse(params)
print(f'retrieval MSE: empty memory {before:.4f} -> after the stream '
f'{after:.4f}')
assert after < before / 5