pytorch_模型參數-保存,加載,打印


1.保存模型參數(gen-我自己的模型名字)

torch.save(self.gen.state_dict(), os.path.join(self.gen_save_path, 'gen_%d.pth'%step)) 

2.加載模型參數

self.gen.load_state_dict(torch.load(os.path.join(self.gen_save_path, 'gen_%d.pth'%step),map_location='cpu'))

3.打印查看模型參數

    pthfile = r'./trained_models\64\models\gen_97000.pth'
    net = torch.load(pthfile,map_location='cpu')
    print(net)

打印結果:

 

 


免責聲明!

本站轉載的文章為個人學習借鑒使用,本站對版權不負任何法律責任。如果侵犯了您的隱私權益,請聯系本站郵箱yoyou2525@163.com刪除。



 
粵ICP備18138465號   © 2018-2025 CODEPRJ.COM