如何实现ZMQ服务器批量处理请求并返回对应客户端响应?
ZMQ批量请求处理的客户端/服务器实现方案
我需要在客户端/服务器模型中使用ZMQ,要求每个服务器接收满100个请求后再进行联合处理,随后向对应的客户端返回这100个响应。这么做是因为服务器的GPU计算只有批量处理时才具备效率。但我尝试的代码触发了zmq.error.ZMQError: Operation cannot be accomplished in current state错误,原因是服务器使用REQ/REP模式时,无法连续接收多个请求——REQ/REP要求严格遵循"接收-发送"的交替流程,不能连续调用recv_pyobj()而不穿插send_pyobj()。
原始错误代码
import multiprocessing as mp import numpy as np import time import zmq def computation(inputs): time.sleep(1) # 模拟GPU计算的固定开销 results = np.zeros((len(inputs), 8)) return results def server(port, batch=100): context = zmq.Context() socket = context.socket(zmq.REP) socket.bind(f'tcp://*:{port}') while True: inputs = np.empty((100, 64)) for i in range(batch): inputs[i] = socket.recv_pyobj() results = computation(inputs) for i in range(batch): socket.send_pyobj(results[i]) def client(ports): context = zmq.Context() socket = context.socket(zmq.REQ) for port in ports: socket.connect(f'tcp://localhost:{port}') while True: input_ = np.zeros(64) socket.send_pyobj(input_) result = socket.recv_pyobj() if __name__ == '__main__': num_clients = 10 num_servers = 3 ports = list(range(5550, 5550 + num_servers)) for port in ports: mp.Process(target=server, args=(port,)).start() for _ in range(num_clients): mp.Process(target=client, args=(ports,)).start()
解决方案:使用ROUTER/DEALER模式
REQ/REP的严格交替限制不适合批量攒请求的场景,改用**ROUTER(服务器端)+ DEALER(客户端)**模式即可解决。这两种套接字支持异步通信:
- ROUTER套接字会自动记录每个客户端的身份标识,每次接收消息时会先收到客户端身份,再收到请求内容
- DEALER套接字允许客户端自由发送请求和接收响应,不需要严格遵循"发送-接收"的交替
修正后的代码
import multiprocessing as mp import numpy as np import time import zmq def computation(inputs): time.sleep(1) # 模拟GPU计算的固定开销 results = np.zeros((len(inputs), 8)) return results def server(port, batch=100): context = zmq.Context() socket = context.socket(zmq.ROUTER) socket.bind(f'tcp://*:{port}') while True: batch_data = [] # 保存(客户端身份, 请求输入)的列表 # 攒够batch个请求 while len(batch_data) < batch: # ROUTER接收顺序: 客户端身份 -> 空帧 -> 请求数据 client_id = socket.recv() socket.recv() # 跳过空帧(ZMQ DEALER发送的消息会带空帧) input_data = socket.recv_pyobj() batch_data.append((client_id, input_data)) # 提取所有输入进行批量计算 inputs = np.array([data[1] for data in batch_data]) results = computation(inputs) # 向每个客户端返回对应的结果 for (client_id, _), result in zip(batch_data, results): # ROUTER发送顺序: 客户端身份 -> 空帧 -> 结果数据 socket.send(client_id, zmq.SNDMORE) socket.send(b'', zmq.SNDMORE) socket.send_pyobj(result) def client(ports): context = zmq.Context() socket = context.socket(zmq.DEALER) # 给每个客户端设置唯一身份(可选,帮助服务器区分) socket.setsockopt_string(zmq.IDENTITY, f'client_{mp.current_process().pid}') for port in ports: socket.connect(f'tcp://localhost:{port}') while True: input_ = np.zeros(64) # DEALER发送消息: 空帧 -> 请求数据(符合ZMQ的协议要求) socket.send(b'', zmq.SNDMORE) socket.send_pyobj(input_) # 接收响应: 空帧 -> 结果数据 socket.recv() # 跳过空帧 result = socket.recv_pyobj() # 可添加结果处理逻辑 # print(f"Client {mp.current_process().pid} received result: {result}") if __name__ == '__main__': num_clients = 10 num_servers = 3 ports = list(range(5550, 5550 + num_servers)) for port in ports: mp.Process(target=server, args=(port,)).start() for _ in range(num_clients): mp.Process(target=client, args=(ports,)).start()
关键改动说明
- 服务器端改用ROUTER套接字:
- 每次接收消息时会先获取客户端的身份标识,确保后续能准确回复对应的客户端
- 可以连续接收任意数量的请求,无需穿插发送操作
- 客户端改用DEALER套接字:
- 解除了REQ模式下"发送必须紧跟接收"的限制,适配服务器的批量处理逻辑
- 发送消息时需要遵循ZMQ协议,先发送空帧再发送实际数据
- 批量处理逻辑:
- 服务器先收集足够数量的请求(包含客户端身份和输入数据),再统一进行GPU计算
- 计算完成后,遍历每个请求对应的客户端身份,逐个发送响应
内容的提问来源于stack exchange,提问作者danijar
相关产品推荐
相关产品推荐

