如何压缩PyTorch模型权重以实现联邦学习场景下的高效网络传输?
解决PyTorch联邦学习模型权重传输体积过大的方案
核心问题分析
pickle/dill序列化PyTorch模型的state_dict时,会额外存储张量的设备信息、梯度状态、Python对象结构等冗余元数据,这是6kb原始权重膨胀到65mb的核心原因。以下是针对性的优化方案:
只传输张量原始二进制数据,剥离冗余元数据
直接提取张量的原始数值字节,仅保留形状、数据类型这类必要元数据,序列化后体积会接近原始权重大小:
import numpy as np import torch # 序列化:仅保留张量核心数据与必要元数据 def serialize_weights(state_dict): serialized = {} for k, v in state_dict.items(): # 转CPU张量→numpy数组→原始字节,同时记录形状和数据类型 tensor_np = v.cpu().numpy() serialized[k] = { 'data': tensor_np.tobytes(), 'shape': tensor_np.shape, 'dtype': str(tensor_np.dtype) } return serialized # 反序列化:从字节恢复张量 def deserialize_weights(serialized, device='cuda'): state_dict = {} for k, info in serialized.items(): arr = np.frombuffer(info['data'], dtype=info['dtype']).reshape(info['shape']) state_dict[k] = torch.tensor(arr, device=device) return state_dict
改用轻量结构化序列化格式替代pickle/dill
msgpack、protobuf这类格式专为结构化数据设计,序列化效率远高于pickle,不会存储Python对象的冗余信息。以msgpack为例:
import msgpack # 序列化:结合numpy张量转码+msgpack打包 def serialize_with_msgpack(state_dict): data = {} for k, v in state_dict.items(): tensor_np = v.cpu().numpy() data[k] = { 'data': tensor_np.tobytes(), 'shape': tensor_np.shape, 'dtype': str(tensor_np.dtype) } return msgpack.packb(data, use_bin_type=True) # 反序列化 def deserialize_with_msgpack(packed_data, device='cuda'): data = msgpack.unpackb(packed_data) state_dict = {} for k, info in data.items(): arr = np.frombuffer(info['data'], dtype=info['dtype']).reshape(info['shape']) state_dict[k] = torch.tensor(arr, device=device) return state_dict
联邦学习专属压缩策略
1. 差值稀疏化传输
仅传输本地权重与全局权重的差值,且只保留差值中的非零元素(适合训练中权重更新稀疏的场景):
def compute_weight_diff(local_state_dict, global_state_dict): diff = {} for k in local_state_dict: diff_k = local_state_dict[k] - global_state_dict[k] non_zero_mask = diff_k != 0 if non_zero_mask.any(): diff[k] = { 'indices': non_zero_mask.nonzero().cpu().numpy().tolist(), 'values': diff_k[non_zero_mask].cpu().numpy().tolist(), 'shape': diff_k.shape, 'dtype': str(diff_k.dtype) } return diff
2. 权重量化
将float32权重量化为float16甚至8位整数,在精度损失可接受的前提下直接将体积减半或压缩到1/4:
def serialize_quantized(state_dict): serialized = {} for k, v in state_dict.items(): # 转float16量化张量 quantized_v = v.half() tensor_np = quantized_v.cpu().numpy() serialized[k] = { 'data': tensor_np.tobytes(), 'shape': tensor_np.shape, 'dtype': str(tensor_np.dtype) } return serialized
多策略组合优化
先做张量轻量化处理(原始字节+量化),再用zlib/gzip等压缩算法二次压缩,实现最大程度的体积缩减:
import zlib def serialize_and_compress(state_dict): # 先量化+原始字节序列化 serialized = serialize_quantized(state_dict) # msgpack打包后zlib压缩 packed = msgpack.packb(serialized, use_bin_type=True) compressed = zlib.compress(packed) return compressed def decompress_and_deserialize(compressed_data, device='cuda'): # 解压缩→msgpack解析→张量恢复 packed = zlib.decompress(compressed_data) serialized = msgpack.unpackb(packed) return deserialize_weights(serialized, device)
内容的提问来源于stack exchange,提问作者ditty
相关产品推荐
相关产品推荐

