如何基于Socket API实现TCP协议下可靠高效的大数据传输?
我正在开展一个客户端-服务端架构下的安全大模型推理项目,需要大量协作计算通信。比如服务端需频繁向客户端发送模型参数(如矩阵),这要求通过Socket实现TCP协议下的可靠大数据传输。直接调用socket.socket.sendall()和socket.socket.recv()的方法不可行,因为TCP是字节流服务,接收方难以区分数据边界。
我编写了如下代码:
import socket import pickle import torch class BetterSocket: def __init__(self, s): self.socket = s self.msg_len = 2 ** 12 def sendall(self, obj): pkl = pickle.dumps(obj) l = len(pkl) init_l = len(pkl) lbs = l.to_bytes(4) print(f"With LBS = {lbs}") self.socket.sendall(lbs) print("Going to send...") while l > self.msg_len: self.socket.sendall(pkl[0:self.msg_len]) pkl = pkl[self.msg_len:] l = l - self.msg_len self.socket.recv(len(b'0')) self.socket.sendall(pkl) print(f"Done. {4 + init_l} bytes have been sent.") def recv(self): print("Waiting to receive...") lbs = self.socket.recv(4) print(f"LBS {lbs} received.") l = int.from_bytes(lbs) print(f"Length {l} received.") pkl = b'' while l > self.msg_len: pkl = pkl + self.socket.recv(self.msg_len) l = l - self.msg_len self.socket.sendall(b'0') pkl = pkl + self.socket.recv(l) obj = pickle.loads(pkl) return obj
该类可传输任意类型对象,每个数据单元头部附加4字节标识数据总大小。循环中的“反向确认”似乎是平衡收发速度的必要机制,否则会抛出_pickle.UnpicklingError: pickle data was truncated错误,表明接收方拆分字节流错误;当self.msg_len设置过大时也会出现相同错误。
但该方案性能过差,传输50000*800的矩阵耗时约5分钟,原因是客户端与服务端频繁进行反向确认。如何在保证传输正确性的前提下,避免如此严重的性能损耗?
优化方案
1. 移除反向确认,确保完整接收数据
TCP本身是可靠传输协议,无需每段数据都发反向确认。之前的截断错误不是因为缺少确认,而是recv()没有保证拿到指定长度的字节——socket.recv(n)最多返回n字节,实际可能更少,必须循环读取直到凑够长度。
先实现一个通用的全量接收函数:
def recv_all(self, size): data = b'' while len(data) < size: # 每次读取剩余需要的字节数和msg_len的较小值,避免内存浪费 chunk = self.socket.recv(min(size - len(data), self.msg_len)) if not chunk: raise ConnectionError("连接已中断") data += chunk return data
然后重构收发方法,去掉反向确认逻辑:
def sendall(self, obj): pkl = pickle.dumps(obj) total_len = len(pkl) # 发送4字节长度头(指定字节序避免跨平台问题) self.socket.sendall(total_len.to_bytes(4, byteorder='big')) # 直接发送完整序列化数据,sendall会确保数据全部发出 self.socket.sendall(pkl) print(f"完成发送,共发送 {4 + total_len} 字节") def recv(self): # 接收长度头 len_data = self.recv_all(4) total_len = int.from_bytes(len_data, byteorder='big') # 接收完整序列化数据 pkl_data = self.recv_all(total_len) return pickle.loads(pkl_data)
2. 增大消息块大小
原代码中msg_len设为4KB(212)太小,建议调整为64KB(216)或128KB(2^17),减少循环次数,提升传输效率。修改__init__方法:
def __init__(self, s): self.socket = s self.msg_len = 2 ** 16 # 64KB,可根据网络带宽调整
3. 替换序列化方式提升效率
pickle对大张量的序列化效率一般,针对PyTorch场景可以用更高效的方式:
- 直接用torch序列化:跳过pickle,专门针对张量优化
import io def sendall(self, tensor): buffer = io.BytesIO() torch.save(tensor, buffer) buffer.seek(0) tensor_data = buffer.getvalue() total_len = len(tensor_data) self.socket.sendall(total_len.to_bytes(4, byteorder='big')) self.socket.sendall(tensor_data) def recv(self): len_data = self.recv_all(4) total_len = int.from_bytes(len_data, byteorder='big') tensor_data = self.recv_all(total_len) buffer = io.BytesIO(tensor_data) return torch.load(buffer) - 用MessagePack压缩序列化:比pickle更快更紧凑,适合跨平台场景
- 添加数据压缩:如果带宽有限,对序列化后的字节流用zstd/gzip压缩,减少传输体积
4. 启用TCP_NODELAY降低延迟
禁用Nagle算法,避免小数据包的合并延迟,适合低延迟需求的大数据传输:
def __init__(self, s): self.socket = s self.socket.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) self.msg_len = 2 ** 16
内容的提问来源于stack exchange,提问作者user24934971

