浏览代码

further fix #90

lchu 1 年之前
父节点
当前提交
80a4c36707
共有 1 个文件被更改,包括 1 次插入1 次删除
  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()