已合并
[Contrib][PyTorch] WaveGlow save model 失败问题 #1228
AtomGit-Bot创建于 2022年7月22日
[Contrib][PyTorch] WaveGlow save model 失败问题 #1228
已合并
从refs/pull/1228/head合入到master
共 1 个文件变更+4-2
| @@ -19,6 +19,8 @@ import argparse | |||
| 19 | import json | 19 | import json |
| 20 | import os | 20 | import os |
| 21 | import torch | 21 | import torch |
| 22 | +if torch.__version__ >= "1.8": | ||
| 23 | + import torch_npu | ||
| 22 | import time | 24 | import 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, iteration | 45 | 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) |