bs,n_keys,n_vals,d = 2,10,30,16
b = dict(
profile=torch.randint(0, 10, (bs, 3, 3)),
profile_mask=torch.ones(bs, 3).bool(),
profile_time=torch.zeros(bs, 3),
lifelong=torch.empty(bs, 0, 3).long(),
lifelong_mask=torch.empty(bs, 0).bool(),
lifelong_time=torch.empty(bs, 0),
event_tokens=torch.randint(0, 10, (5, 3)),
event_offsets=torch.tensor([0, 2, 5]),
event_user=torch.tensor([0, 1]),
event_time=torch.tensor([1., 2.]),
cal=torch.tensor([[12, 2, 15], [18, 5, 20]]),
history_offsets=torch.tensor([0, 1, 2]),
uids=['u1', 'u2'],
event_labels=torch.randint(0, n_vals, (5,)),
mlm_mask=torch.tensor([1, 0, 1, 0, 1]).bool())
m = PRAGMAModel(n_keys, n_vals, d_model=d, n_heads=4, prof_layers=1, event_layers=1, hist_layers=1, p=0.)
res = m(b)
test_eq(res['h_usr'].shape, (bs, d))
test_eq(res['h_evt'].shape, (2, d))
test_eq(res['logits'].shape, (3, n_vals))
test_eq(res['labels'].shape, (3,))
test_eq(res['loss'].ndim, 0)
torch.isfinite(res['loss'])
res['loss']