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

联邦学习中基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 00:33:37