如何通过Socket发送多维NumPy数组列表并还原格式?
通过Socket传输多维NumPy数组列表并还原格式的解决方案
问题背景
需要通过Socket将包含多维NumPy数组的列表发送至服务器,并在接收后还原原始格式。数组列表示例如下:
[array([[[[-1.04182057e-01, 9.81570184e-02, 8.69736895e-02, -6.61955923e-02, -4.51700203e-02], [ 5.26290983e-02, -1.18473642e-01, 2.64136307e-02, -9.26332623e-02, -6.63961545e-02], [-8.80082026e-02, 7.90973455e-02, -1.13944486e-02, -1.51292123e-02, 7.65037686e-02], [-9.15177837e-02, 7.08795676e-04, -1.08281896e-03, 8.65678713e-02, 6.68114647e-02], [-8.45356733e-02, -6.90313280e-02, -5.81113175e-02, -1.14920050e-01, -4.11906727e-02]], ... 3.35839503e-02, 6.30911887e-02, 4.10411768e-02, -3.64055522e-02, -3.56383622e-02, 9.80690420e-02, 8.15757737e-02, -1.00057133e-01, 1.16158882e-02, -9.82330441e-02, 9.00610462e-02, -1.01473713e-02, -2.64037345e-02, 1.37711661e-02, 6.63968623e-02]], dtype=float32), array([-0.02089943, -0.0020895 , -0.00506333, 0.03931976, 0.04795408, -0.01520141, -0.03287903, 0.0037387 , 0.01339047, -0.0576841 ], dtype=float32)]
此前尝试json.dumps发送时触发TypeError: Object of type ndarray is not JSON serializable;直接转字符串编码发送后,服务器仅能收到普通字符串,无法还原为多维NumPy数组列表。使用Python 3.10版本。
解决方案:用Pickle序列化二进制传输
核心思路是利用Python原生的pickle模块序列化NumPy对象,直接传输二进制数据,同时添加数据长度标记避免Socket粘包问题。
客户端代码
import socket import pickle import numpy as np from typing import List # 假设parameters_to_ndarrays为已定义的转换函数 aggregated_ndarrays: List[np.ndarray] = parameters_to_ndarrays(aggregated_parameters) n_examples_fit = 1000 # 替换为实际样本数 print("Attempting server connection") conn = socket.socket(socket.AF_INET, socket.SOCK_STREAM) conn.connect(("127.0.0.1", 8088)) # 打包并序列化数据 data_package = (aggregated_ndarrays, n_examples_fit) serialized_data = pickle.dumps(data_package) # 先发送数据长度(4字节大端存储) conn.sendall(len(serialized_data).to_bytes(4, byteorder='big')) # 再发送序列化后的二进制数据 conn.sendall(serialized_data) conn.close()
服务器代码
import socket import pickle import numpy as np sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) sock.bind(("127.0.0.1", 8088)) sock.listen() print("Created server socket and listening %s" % sock) conn, addr = sock.accept() print("Accepted client connection from", addr) # 接收数据长度 length_bytes = conn.recv(4) data_length = int.from_bytes(length_bytes, byteorder='big') # 接收完整序列化数据 received_data = b'' while len(received_data) < data_length: chunk = conn.recv(min(data_length - len(received_data), 4096)) if not chunk: break received_data += chunk # 反序列化还原数据 aggregated_ndarrays, n_examples_fit = pickle.loads(received_data) # 验证结果 print("Received training examples count:", n_examples_fit) print("First array shape:", aggregated_ndarrays[0].shape) print("First array dtype:", aggregated_ndarrays[0].dtype) conn.close() sock.close()
关键说明
- Pickle的优势:原生支持NumPy数组等复杂Python对象的序列化,自动保留数组的维度、数据类型等所有元信息,无需额外处理。
- 粘包处理:Socket传输可能出现数据粘包,先发送固定长度的数据长度标记,服务器据此接收完整的二进制数据,避免截断或数据合并。
- 安全提示:Pickle反序列化不可信数据存在安全风险,若服务对外开放,建议搭配身份验证,或改用
msgpack等更安全的序列化方案并显式处理类型。
内容的提问来源于stack exchange,提问作者kyro121
相关产品推荐
相关产品推荐

