联邦学习中基于Python线程实现模型平均的互斥同步问询
问题描述
在实现联邦学习算法时,通过以下Python代码创建客户端线程:
sockets_thread = [] no_of_client = 1 all_data = b"" while True: try: for i in range(no_of_client): connection, client_info = soc.accept() print("\nNew Connection from {client_info}.".format(client_info=client_info)) socket_thread = SocketThread(connection=connection, client_info=client_info, buffer_size=1024, recv_timeout=100) sockets_thread.append(socket_thread) for i in range(no_of_client): sockets_thread[i].start() sockets_thread[i].join() except: soc.close() print("(Timeout) Socket Closed Because no Connections Received.\n") break
SocketThread类的run及reply方法代码如下:
class SocketThread(object): def run(self): while True: received_data, status = self.recv() if status == 0: self.connection.close() break self.reply(received_data) def reply(self, received_data): model = SimpleASR() #all threads must averge the model before going to next line model_instance = self.model_averaging(model, model_instance) print("All threads completed model averging.") #now do rest of the things
要求model_instance = self.model_averaging(model, model_instance)函数执行时需互斥,且所有线程完成模型平均后才能继续执行后续代码,需使用Python条件变量实现。
解决方案
要实现线程间的互斥执行和同步等待,需结合threading.Lock(保证互斥)与threading.Condition(实现条件等待),同时用全局计数器跟踪已完成模型平均的线程数量。
核心实现步骤
- 共享状态初始化:创建所有线程共用的锁、条件变量、线程计数器和目标线程数。
- 修改SocketThread类:让类接收共享的同步组件作为初始化参数,确保线程间共享状态。
- 互斥与同步逻辑:在模型平均代码块外包裹条件变量上下文,完成后更新计数器,根据计数器状态决定唤醒等待线程或进入等待。
完整代码示例
import threading # 全局共享同步组件,所有SocketThread实例共用 lock = threading.Lock() condition = threading.Condition(lock) completed_threads = 0 target_threads = 1 # 与原代码no_of_client保持一致 class SocketThread(object): def __init__(self, connection, client_info, buffer_size, recv_timeout): self.connection = connection self.client_info = client_info self.buffer_size = buffer_size self.recv_timeout = recv_timeout # 注入共享同步组件 self.condition = condition def run(self): while True: received_data, status = self.recv() if status == 0: self.connection.close() break self.reply(received_data) def reply(self, received_data): model = SimpleASR() global completed_threads with self.condition: # 互斥执行模型平均 model_instance = self.model_averaging(model, model_instance) # 更新已完成线程计数 completed_threads += 1 # 检查是否所有线程完成模型平均 if completed_threads == target_threads: # 唤醒所有等待的线程 self.condition.notify_all() # 重置计数器,为下一轮联邦学习准备 completed_threads = 0 else: # 未全部完成则等待通知 while completed_threads < target_threads: self.condition.wait() # 所有线程完成后执行后续逻辑 print("All threads completed model averaging.") # 后续业务代码... def recv(self): # 补充原有recv方法实现,此处为模拟返回 return b"dummy_data", 1 def model_averaging(self, model, model_instance): # 补充模型平均的具体实现 return model # 原有线程创建逻辑 sockets_thread = [] no_of_client = target_threads all_data = b"" while True: try: for i in range(no_of_client): connection, client_info = soc.accept() print("\nNew Connection from {client_info}.".format(client_info=client_info)) socket_thread = SocketThread(connection=connection, client_info=client_info, buffer_size=1024, recv_timeout=100) sockets_thread.append(socket_thread) for i in range(no_of_client): sockets_thread[i].start() sockets_thread[i].join() except: soc.close() print("(Timeout) Socket Closed Because no Connections Received.\n") break
关键逻辑说明
- 互斥性保障:
with self.condition会自动获取底层锁,确保同一时间只有一个线程执行model_averaging。 - 同步等待机制:线程完成模型平均后,若未达到目标线程数则调用
wait()释放锁并进入等待;最后一个线程完成时调用notify_all()唤醒所有等待线程。 - 计数器重置:所有线程完成后重置计数器,避免影响下一轮联邦学习的同步逻辑。
内容的提问来源于stack exchange,提问作者treeoid
相关产品推荐
相关产品推荐

