From e9721154ed94dbedb3443bac83a8e1e3f9d4b734 Mon Sep 17 00:00:00 2001 From: JoJoBarthold2 Date: Tue, 5 Sep 2023 10:25:28 +0200 Subject: model now saves validation_loss and cer and wer --- swr2_asr/train.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) (limited to 'swr2_asr/train.py') diff --git a/swr2_asr/train.py b/swr2_asr/train.py index 63deb72..6f3bc6c 100644 --- a/swr2_asr/train.py +++ b/swr2_asr/train.py @@ -253,7 +253,7 @@ def run( iter_meter, ) - test( + test_loss,avg_cer,avg_wer = test( model=model, device=device, test_loader=valid_loader, @@ -262,7 +262,12 @@ def run( ) print("saving epoch", str(epoch)) torch.save( - {"epoch": epoch, "model_state_dict": model.state_dict(), "loss": loss}, + {"epoch": epoch, + "model_state_dict": model.state_dict(), + "loss": loss, + "test_loss": test_loss, + "avg_cer": avg_cer, + "avg_wer": avg_wer}, path + str(epoch), ) -- cgit v1.2.3