INNER CODE UNIT · Python

n_eval_iters

Guitaricet/relora · torchrun_main.py:160

            n_eval_iters = int(target_eval_tokens / tokens_in_batch_info[0])

        if target_eval_tokens != -1 and i > n_eval_iters: break

        batch = {k: v.to(device) for k, v in batch.items()}

        loss = model(**batch, labels=batch["input_ids"]).loss
        if torch.isnan(ddp_loss_info[0]):
            print(f"Rank {dist.get_rank()} got nan loss. This is probably a bug.")

        tokens_in_batch = batch["input_ids"].numel()
        assert tokens_in_batch > 0, "Batch size is zero"
        ddp_loss_info[0] += loss.detach()
        ddp_loss_info[1] += 1
        ddp_loss_info[2] += tokens_in_batch

    # check if loss is nan
    if torch.isnan(ddp_loss_info[0]):

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…