class Client(object):
    # 构造函数
    def __init__(self, conf, model, train_dataset, id = 1):
        # 配置文件
        self.conf = conf
        # 客户端本地模型(一般由服务器传输)
        self.local_model = model
        # 客户端ID
        self.client_id = id
        # 客户端本地数据集
        self.train_dataset = train_dataset
        # 用来划分训练数据集的。首先,将训练数据集的索引全部存储在一个列表`all_range‘中。然后,将训练数据集划分为`no_models`个子集,每个子集的大小为`data_len`。
        # 其中,`id` 表示当前模型的编号,从0开始。通过`all_range[id * data_len: (id + 1) * data_len]`可以获取当前模型对应的训练数据集的索引列表。
        # 最后,使用 `torch.utils.data.DataLoader`来生成一个数据加载器,其中`sampler`参数指定了使用`SubsetRandomSampler`采样器来从训练数据集的子集中随机采样数据。
        # `batch_size`参数指定了每个batch加载多少个样本。
        all_range = list(range(len(self.train_dataset)))
        data_len = int(len(self.train_dataset) / self.conf['no_models'])
        indices = all_range[id * data_len: (id + 1) * data_len]
        # 生成一个数据加载器
        self.train_loader = torch.utils.data.DataLoader(
            # 制定父集合
            self.train_dataset,
            # batch_size每个batch加载多少个样本(默认: 1)
            batch_size=conf['batch_size'],
            sampler=torch.utils.data.sampler.SubsetRandomSampler(indices)
        )

        # 存储所有训练轮次的平均梯度L2范数
        self.grad_norms = []

        # 选取历史梯度L2范数的p-百分位数
    def get_percentile(self, grad_norm, percentile):
            return np.percentile(grad_norm, percentile)

    #     # 计算裁剪阈值
    # def get_clip_threshold(grad_norm, percentile):
    #         percentile_value = get_percentile(grad_norm, percentile)
    #         return percentile_value


        # 模型本地训练函数
    def local_train(self, model):
        # state_dict()方法返回模型的所有参数及其对应的名称。for循环中的name变量表示参数名称,param变量表示参数值
        # items()`是Python字典(dictionary)的一个方法,用于返回字典中所有键值对(key-value pairs)的元组(tuple)列表。在这里,
        # `model.state_dict()`返回的是一个字典,其中每个键(key)是模型的一个参数的名称,对应的值(value)是该参数的值
        for name, param in model.state_dict().items():
            # 客户端首先用服务器端下发的全局模型覆盖本地模型
            self.local_model.state_dict()[name].copy_(param.clone())
            # print('name, param:下发下的数据', name, param)
        # 定义最优化函数器用于本地模型训练
        optimizer = torch.optim.SGD(self.local_model.parameters(), lr=self.conf['lr'], momentum=self.conf['momentum'])

        # 本地训练模型,将模型设置为训练模式,即启用BatchNormalization和Dropout等特定于训练的操作
        self.local_model.train()        # 设置开启模型训练(可以更改参数)



        # 开始训练模型
        for e in range(1, self.conf['local_epochs'] + 1):
            print('本地训练的第几轮:', e)
            local_grads = []
            # 计算每一个数据的梯度
            for batch_id, batch in enumerate(self.train_loader):
                # print('batch_id, batch', batch_id, batch)
                # print('self.train_loader', len(self.train_loader))
                # `data`和`target`是在训练过程中一个batch的输入数据和标签,通常情况下,`data`是一个张量,包含了一组图像数据,而`target`是一个张量,包含了这些图像数据对应的标签。在深度学习中,我们通常将输入数据和标签一起组成一个batch进行训练,以提高训练效率。
                data, target = batch
                # 加载到gpu
                if torch.cuda.is_available():
                    data = data.cuda()
                    target = target.cuda()
                # 将梯度设置为0,进行前向传播和反向传播计算梯度,并根据计算出的梯度更新模型参数。
                optimizer.zero_grad()
                # 训练预测,本地模型对数据 `data` 进行前向传播得到预测结果的操作。即模型对数据的预测结果。
                output = self.local_model(data)
                # 计算损失函数 cross_entropy交叉熵误差
                loss = torch.nn.functional.cross_entropy(output, target)
                # 反向传播,计算梯度
                loss.backward()
                optimizer.step()



            # 计算梯度的L2范数并相加
            total_norm = 0
            # 返回模型中所有需要训练的参数,这些参数是可以被优化器进行更新的
            for param in model.parameters():
                # print('梯度:', param.grad)
                if param.grad is not None:
                    # 计算梯度张量的L2范数
                    param_norm = param.grad.data.norm(2)
                    # 将计算结果转换为Python数字,然后将其平方并添加到total_norm中。最后,total_norm的平方根将用于进行梯度裁剪。
                    # 如果张量只包含一个元素,则使用`item()`方法可将该元素转换为Python数字。如果张量中有多个元素,则不能使用`item()`方法。在这种情况下,需要使用其他方法,如`tolist()`或`numpy()`来转换张量。
                    total_norm += param_norm.item() ** 2
            # 计算平均梯度的L2范数
            avg_norm = total_norm ** 0.5 / self.conf['batch_size']
            # 输出平均梯度的L2范数
            print('Average gradient L2 norm:', avg_norm)
            self.grad_norms.append(avg_norm)
            clip_threshold = self.get_percentile(self.grad_norms, self.conf['percentile'])
            print('grad', self.grad_norms)
            print('clip_threshold:', clip_threshold)


            # 梯度裁减
            # diff = dict()
            avg_grads = []
            grads = []
            for params in model.parameters():
                if params.grad is not None:
                        # 使用clip_threshold对模型的参数梯度进行裁剪,以避免梯度爆炸的问题。具体来说,它计算模型所有参数的梯度的L2范数,
                        # 并将其与clip_threshold进行比较。如果梯度的L2范数超过了clip_threshold,那么就将梯度进行缩放,使其L2范数等于L2除以clip_threshold。
                        # 如果梯度的L2范数小于等于clip_threshold,那么就不做任何处理。最后,该函数返回梯度的L2范数。
                        # print('梯度数据', params.grad)
                        # print('L2梯度', params.grad.data.norm(2))
                        # total_norm2+=torch_utils.clip_grad_norm_(model.parameters(), clip_threshold)
                        # print('原梯度', params.grad)
                        # 计算梯度张量的L2范数
                        params_norm = params.grad.data.norm(2)
                        # 使用 `torch.clamp` 函数对梯度进行裁剪,以防止梯度爆炸
                        params.grad.data.clamp_(-clip_threshold, clip_threshold)
                        grads.append(params.grad.data.clone())
            local_grads.append(grads)
            for i in range(len(local_grads[0])):
                            avg_grad = torch.zeros_like(local_grads[0][i])
                            for j in range(len(local_grads)):
                                avg_grad += local_grads[j][i]
                            # 将平均梯度除以参与计算平均梯度的模型数量
                            avg_grad /= len(local_grads)
                            avg_grads.append(avg_grad)
                            print('平均', avg_grads)
            for i, param in enumerate(self.local_model.parameters()):
                            new_data = param.data - 0.001 * avg_grads[i]
                            param.data = new_data
            return param.data, clip_threshold

原文地址: https://www.cveoy.top/t/topic/nx1N 著作权归作者所有。请勿转载和采集!

免费AI点我,无需注册和登录