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):

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…