已合并
[Contrib][PyTorch] WaveGlow save model 失败问题 #1228
AtomGit-Bot创建于 2022年7月22日
[Contrib][PyTorch] WaveGlow save model 失败问题 #1228
已合并
AtomGit-Bot创建于 2022年7月22日
refs/pull/1228/head合入到master
1 个文件变更+4-2
MPyTorch/contrib/audio/WaveGlow/train.py+4-2
@@ -19,6 +19,8 @@ import argparse
19import json19import json
20import os20import os
21import torch21import torch
22+if torch.__version__ >= "1.8":
23+ import torch_npu
22import time24import time
23 25 
24#=====START: ADDED FOR DISTRIBUTED======26#=====START: ADDED FOR DISTRIBUTED======
@@ -37,7 +39,7 @@ def load_checkpoint(checkpoint_path, model, optimizer):
37 iteration = checkpoint_dict['iteration']39 iteration = checkpoint_dict['iteration']
38 optimizer.load_state_dict(checkpoint_dict['optimizer'])40 optimizer.load_state_dict(checkpoint_dict['optimizer'])
39 model_for_loading = checkpoint_dict['model']41 model_for_loading = checkpoint_dict['model']
40- model.load_state_dict(model_for_loading.state_dict())42+ model.load_state_dict(model_for_loading)
41 print("Loaded checkpoint '{}' (iteration {})" .format(43 print("Loaded checkpoint '{}' (iteration {})" .format(
42 checkpoint_path, iteration))44 checkpoint_path, iteration))
43 return model, optimizer, iteration45 return model, optimizer, iteration
@@ -47,7 +49,7 @@ def save_checkpoint(model, optimizer, learning_rate, iteration, filepath):
47 iteration, filepath))49 iteration, filepath))
48 model_for_saving = WaveGlow(**waveglow_config).to("npu:0")50 model_for_saving = WaveGlow(**waveglow_config).to("npu:0")
49 model_for_saving.load_state_dict(model.state_dict())51 model_for_saving.load_state_dict(model.state_dict())
50- torch.save({'model': model_for_saving,52+ torch.save({'model': model_for_saving.state_dict(),
51 'iteration': iteration,53 'iteration': iteration,
52 'optimizer': optimizer.state_dict(),54 'optimizer': optimizer.state_dict(),
53 'learning_rate': learning_rate}, filepath)55 'learning_rate': learning_rate}, filepath)