MCPcopy Create free account
hub / github.com/Pints-AI/1.5-Pints / validate

Function validate

pretrain/main.py:642–669  ·  view source on GitHub ↗

Run validation and calculate loss.

(
    fabric: lightning.Fabric,
    model: torch.nn.Module,
    val_dataloader: DataLoader,
    eval_iters: int,
)

Source from the content-addressed store, hash-verified

640
641@torch.no_grad()
642def validate(
643 fabric: lightning.Fabric,
644 model: torch.nn.Module,
645 val_dataloader: DataLoader,
646 eval_iters: int,
647) -> torch.Tensor:
648 """Run validation and calculate loss."""
649
650 fabric.print('Validating ...')
651 model.eval()
652
653 losses = torch.zeros(eval_iters, device=fabric.device)
654 for k, val_data in enumerate(val_dataloader):
655 if k >= eval_iters:
656 break
657 input_ids = val_data[:, 0 : model.config.block_size].contiguous()
658 targets = val_data[:, 1 : model.config.block_size + 1].contiguous()
659 logits = model(input_ids)
660 loss = chunked_cross_entropy(logits, targets, chunk_size=0)
661
662 # loss_func = FusedCrossEntropyLoss()
663 # loss = loss_func(logits, targets)
664 losses[k] = loss.item()
665
666 out = losses.mean()
667
668 model.train()
669 return out
670
671
672def create_dataloader(

Callers 1

trainFunction · 0.70

Calls 1

chunked_cross_entropyFunction · 0.90

Tested by

no test coverage detected