INNER CODE UNIT · Python
evaluate
lucidrains/x-transformers · train_discover_reversal.py:101
def evaluate():
model.eval()
accs = {}
for length in range(1, len_max + 1):
tokens = torch.randint(0, num_fillers, (256, length), device = device)
lengths = torch.full((256,), length, device = device)
outputs = model.decode(model.encode(tokens), lengths)
accs[length] = (outputs[:, :length] == tokens.flip(-1)).all(dim = 1).float().mean().item()
return accs
print(f'training bottleneck decoder on reversal (lists of length 1..{len_max}, letters A..Z) [{device}]')
best_acc = -1.
best_state = None
for i in range(train_steps):