You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将PyTorch神经网络模型JSON编码用于联邦学习节点交互?

PyTorch模型JSON序列化解决方案(联邦学习场景)

问题根源

PyTorch的nn.Module实例包含大量不可JSON序列化的对象(如张量、C++后端绑定模块),直接用json.dumps()序列化必然失败。联邦学习场景下,无需传输完整模型对象,只需传递模型权重参数或结构化的模型定义+权重即可。

先修复你的模型代码bug

你的代码存在两处致命问题,会导致模型无法正常运行:

  1. 定义了两个forward方法,后一个会覆盖前一个,且后一个中使用的fc1、fc2、output在__init__中未定义;
  2. 第一个forward方法没有返回值。

修复后的核心代码片段:

import torch
import torch.nn as nn
import torch.nn.functional as F

class Net(nn.Module):  # 类名建议大写开头,符合PEP8规范
    def __init__(self):
        super(Net, self).__init__()
        self.c1 = nn.Conv2d(3, 32, 3)
        self.b1 = nn.BatchNorm2d(32)
        self.c2 = nn.Conv2d(32, 64, 3)
        self.p1 = nn.MaxPool2d(2, 2)
        self.c3 = nn.Conv2d(64, 128, 3)
        self.d1 = nn.Dropout(0.5)
        self.c4 = nn.Conv2d(128, 256, 5)
        self.p2 = nn.MaxPool2d(2, 2)
        self.c5 = nn.Conv2d(256, 128, 3)
        self.p3 = nn.MaxPool2d(2, 2)
        self.d2 = nn.Dropout(0.2)
        self.l1 = nn.Linear(29*29*128, 128)
        self.l2 = nn.Linear(128, 32)
        self.l3 = nn.Linear(32 , 10)

        self.weights_initialization()

    def forward(self, x):
        x = F.relu(self.c1(x))
        x = self.b1(x)
        x = F.relu(self.c2(x))
        x = self.p1(x)
        x = F.relu(self.c3(x))
        x = self.d1(x)
        x = F.relu(self.c4(x))
        x = self.p2(x)
        x = F.relu(self.c5(x))
        x = self.p3(x)
        x = self.d2(x)
        x = x.view(x.size(0), -1)
        x = F.relu(self.l1(x))
        x = self.d2(x)
        x = F.relu(self.l2(x))
        return self.l3(x)  # 添加返回值

    def weights_initialization(self):
        for m in self.modules():
            if isinstance(m, nn.Linear):
                nn.init.xavier_normal_(m.weight)
                nn.init.constant_(m.bias, 0)

model = Net()

解决方案1:传输模型权重(联邦学习首选)

联邦学习中通常只需要传递模型的权重参数(state_dict),将其转换为JSON可序列化的格式:

序列化权重

import json

def serialize_model_weights(model):
    # 将state_dict中的张量转为列表
    state_dict = model.state_dict()
    serializable_dict = {k: v.tolist() for k, v in state_dict.items()}
    # 转为JSON字符串
    json_str = json.dumps(serializable_dict)
    return json_str

# 使用示例
json_weights = serialize_model_weights(model)

反序列化权重

def deserialize_model_weights(json_str, model):
    # 解析JSON
    serializable_dict = json.loads(json_str)
    # 将列表转回张量
    state_dict = {k: torch.tensor(v) for k, v in serializable_dict.items()}
    # 加载到模型
    model.load_state_dict(state_dict)
    return model

# 使用示例
new_model = Net()
new_model = deserialize_model_weights(json_weights, new_model)

解决方案2:自定义JSON编码器(仅作参考,不推荐)

如果需要序列化模型结构+权重,可以自定义JSON编码器处理张量:

class TorchJSONEncoder(json.JSONEncoder):
    def default(self, obj):
        if isinstance(obj, torch.Tensor):
            return obj.tolist()
        elif isinstance(obj, nn.Module):
            return {
                "state_dict": {k: v.tolist() for k, v in obj.state_dict().items()},
                "model_structure": str(obj)
            }
        return super().default(obj)

# 序列化
json_str = json.dumps(model, cls=TorchJSONEncoder)

# 反序列化需手动重建模型结构再加载权重,效率远低于方案1

注意事项

  • 联邦学习优先用方案1,仅传输权重数据量更小,更适合网络传输;
  • 确保worker节点使用相同的模型结构,应将模型定义代码同步到所有节点,而非序列化整个模型对象;
  • 若追求更高效率,可将torch.save生成的二进制数据转为Base64字符串嵌入JSON,比纯JSON传输权重更快。

内容的提问来源于stack exchange,提问作者Giacomo Di Vaira

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.12 07:45:49