lchu 1 gadu atpakaļ
vecāks
revīzija
80a4c36707
1 mainītis faili ar 1 papildinājumiem un 1 dzēšanām
  1. 1 1
      utils/train_utils.py

+ 1 - 1
utils/train_utils.py

@@ -135,7 +135,7 @@ def train(model, train_dataloader,eval_dataloader, tokenizer, optimizer, lr_sche
         lr_scheduler.step()
           
         if train_config.run_validation:
-            eval_ppl, eval_epoch_loss = evaluation(model, train_config, eval_dataloader, rank, tokenizer)   
+            eval_ppl, eval_epoch_loss = evaluation(model, train_config, eval_dataloader, local_rank, tokenizer)
             if train_config.save_model and eval_epoch_loss < best_val_loss:
                 if train_config.enable_fsdp:
                     dist.barrier()