INNER CODE UNIT · Python

evaluate_model

Guitaricet/relora · torchrun_main.py:144

def evaluate_model(model: nn.Module, eval_dataloader, device, target_eval_tokens=10_000_000):
    _time = time.time()
    was_training = model.train
    model.eval()

    ddp_loss_info = torch.zeros(3).to(device)  # [loss, n_batches, n_tokens]
    tokens_in_batch_info = torch.zeros(1).to(device)

    rank = dist.get_rank()
    for i, batch in enumerate(eval_dataloader):
        if i == 0:
            # this way of estiming the number of eval steps
            # is needed to avoid a deadlock when using FSDP
            batch["input_ids"]: torch.Tensor
            tokens_in_batch_info[0] += batch["input_ids"].numel()
            dist.all_reduce(tokens_in_batch_info, op=dist.ReduceOp.SUM)
            n_eval_iters = int(target_eval_tokens / tokens_in_batch_info[0])

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…