pytorch & tensorflow 保存和加载模型

1. Pytorch

1.1.1 save网络结构和参数:

注意最后一行为"self.state_dict()"

python 复制代码
    def save(self,t):
        current_path = os.path.dirname(os.path.abspath(__file__))
        model_path = 'model/2E_model_' + t + '_'+self.name+'/'

        save_path = os.path.join(current_path,model_path)
        if not os.path.exists(save_path):
            os.makedirs(save_path)

        save_file_path=os.path.join(save_path, 'model.pth')

        torch.save(self.state_dict(),save_file_path)

1.1.2 对应的加载模型参数:

注意对应"agent.load_state_dict(checkpoint)"

python 复制代码
    def load(self,agent,model_path):
        model_pth = 'model.pth'
        model_path = os.path.join(model_path,model_pth)
        checkpoint = torch.load(model_path)
        agent.load_state_dict(checkpoint)
        agent.eval()

1.2.1 保存整个模型

注意为"torch.save(self.model,save_file_path)"

python 复制代码
    def save(self,t):
        current_path = os.path.dirname(os.path.abspath(__file__))
        model_path = 'model/model_' + t + '_'+self.name+'/'

        save_path = os.path.join(current_path,model_path)
        if not os.path.exists(save_path):
            os.makedirs(save_path)

        save_file_path=os.path.join(save_path, 'model.pth')

        torch.save(self.model,save_file_path)

1.2.2 加载整个模型

注意"self.model = torch.load(model_path)"

python 复制代码
    def load(self,model_path):
        model_pth = 'model.pth'
        model_path = os.path.join(model_path,model_pth)
        self.model = torch.load(model_path)
        self.model.eval()

如果没对应上会报错:torch.nn.modules.module.ModuleAttributeError: object has no attribute 'copy',参考此链接

2. Tensorflow

2.1 保存模型

python 复制代码
    def save(self,time):
        current_path = os.path.dirname(os.path.abspath(__file__))
        model_path='model/model_'+time+'_'+self.name+'/weights_'+self.name
        save_path = os.path.join(current_path,model_path)
        if not os.path.exists(save_path):os.makedirs(save_path)
        self.saver.save(self.sess,save_path)

2.2 加载模型

python 复制代码
    def load(self,model_path):
        meta_path = 'weights_'+self.name+'.meta'

        mata_path_dir = os.path.join(model_path,meta_path)

        self.saver = tf.compat.v1.train.import_meta_graph(mata_path_dir)
        a=model_path+'/'
        self.saver.restore(self.sess, tf.train.latest_checkpoint(a))
相关推荐
学着改变2753 分钟前
2026便携式超声波流量计性能白皮书 户外巡检适用性横评
大数据·网络·人工智能·科技·产品运营·量子计算
CVer儿6 分钟前
cuda的跨线程设计sdk方法和需要注意点
人工智能
盼小辉丶7 分钟前
PyTorch强化学习实战——融合人类示范数据的高效强化学习
人工智能·pytorch·python·深度学习·强化学习
空堂与归13 分钟前
序列数据怎么喂给神经网络?用循环神经网络RNN拆开一个时间步
人工智能·rnn·神经网络·nlp
尺度商业17 分钟前
大变革下的电力新序章:供需重塑、错配破解与价值重估
人工智能
天远Date Lab20 分钟前
零信任架构实战:基于天远车信盟出险构建自动化汽车消费贷合规网关
人工智能·机器学习·计算机视觉·ocr
知见漫记23 分钟前
AI文档总结折叠手机推荐,联想moto razr Fold让信息处理更高效
人工智能·智能手机
虎虎(_ _)。゜zzZ26 分钟前
Qdrant向量数据库工程实战
数据库·人工智能·大模型·向量数据库·rag·qdrant
DM今天肝到几点?27 分钟前
GPT-6 Astra 发布:ARC-AGI-3 从 7.8% 跳到 99.9%,OpenAI 宣告「欢迎进入 AGI 时代」
人工智能·gpt·深度学习·agi
智途 Tech33 分钟前
2026年分析多个Excel、CSV和网页数据的AI工具清单:Tabbit 浏览器多源引用
大数据·人工智能·excel