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

如何基于Socket API实现TCP协议下可靠高效的大数据传输?

大模型推理项目中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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 13:58:38