INNER CODE UNIT · Python

loss

lucidrains/x-transformers · train_discover_reversal.py:126

        loss = torch.nn.functional.cross_entropy(logits.reshape(-1, num_tokens), tgt.reshape(-1))

        loss.backward()
        optimizer.step()
        optimizer.zero_grad()

        if (i + 1) % 1000 == 0:
            accs = evaluate()

            if sum(accs.values()) > best_acc:
                best_acc = sum(accs.values())
                best_state = {name: param.detach().clone() for name, param in model.named_parameters()}

            print(f'step {i + 1}: loss {loss.item():.3f} | accuracy by length ' + ' '.join(f'len {l}: {a:.3f}' for l, a in accs.items()), flush = True)

    if exists(best_state):
        model.load_state_dict(best_state)

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…