基于RFF的联邦在线学习KLMS模型MSE异常问题问询
复现Gogineni等人2022论文时的MSE不收敛问题
实现概要
- 客户端数量:100
- 全局迭代次数:1000
- RFF维度:200
- 学习率:0.75
- 每次迭代参与客户端数:20
- 独立蒙特卡洛(Monte Carlo)试验次数:500
每次全局迭代流程:选取部分客户端,各客户端用流数据更新本地模型,上传模型更新至服务器,服务器聚合得到新全局模型。
问题描述
仿真中计算的均方误差(MSE)未按预期收敛或下降,反而大幅波动,无学习过程应有的稳定下降趋势,与论文仿真结果不符。
仿真关键环节
- 客户端输入信号:一阶自回归(AR)模型生成,参数从均匀分布采样(符合论文要求)
- 本地非线性回归:基于RFF的KLMS算法
- 全局模型聚合:对每次选中客户端的权重更新做迭代平均
代码片段
import numpy as np import matplotlib.pyplot as plt # Hyperparameters num_clients = 100 # Number of clients in the simulation independent_experiment = 10 # Number of independent Monte Carlo trials feature_dim = 5 # Dimensionality of the input features rff_dim = 200 # Dimensionality of the random Fourier features num_participating_clients = 20 # Number of clients participating in each iteration learning_rate = 0.75 # Learning rate for the local updates num_iterations = 1000 # Number of iterations for training # Initialize an array to store the MSE values across all trials mse_values_all_trials = np.zeros(num_iterations) # Main loop for averaging over multiple Monte Carlo trials for _ in range(independent_experiment): global_weights = np.zeros(rff_dim) # Initialize global weights x = np.zeros((num_clients, num_iterations, feature_dim)) # Input features for each client y = np.zeros((num_clients, num_iterations, 1)) # Target values for each client z = np.zeros((num_clients, num_iterations, rff_dim)) # Random Fourier features for each client W = np.random.randn(num_clients, feature_dim, rff_dim) # Random weights for RFF b = np.random.uniform(0, 2 * np.pi, (num_clients, 1, rff_dim)) # Random bias for RFF # Generate data for each client for k in range(num_clients): theta_k = np.random.uniform(0.2, 0.9) # Autoregressive coefficient mu_k = np.random.uniform(-0.2, 0.2) # Mean of the process noise sigma2_uk = np.random.uniform(0.2, 1.2) # Variance of the process noise sigma2_nuk = np.random.uniform(0.005, 0.03) # Variance of the observation noise uk = np.random.normal(mu_k, np.sqrt(sigma2_uk), (num_iterations, feature_dim)) # Process noise nuk = np.random.normal(0, np.sqrt(sigma2_nuk), (num_iterations, 1)) # Observation noise # Generate the time series data x[k, 0] = uk[0] for n in range(1, num_iterations): x[k, n, :] = theta_k * x[k, n-1, :] + np.sqrt(1 - theta_k**2) * uk[n] y[k, n, :] = (np.sqrt(x[k, n, 0]**2 + np.sin(np.pi * x[k, n, 3])**2) + (0.8 - 0.5*np.exp(-x[k, n, 1]**2)*x[k, n, 2])) + nuk[n] # Compute the random Fourier features z[k, :, :] = np.sqrt(2 / rff_dim) * np.cos(np.dot(x[k, :, :], W[k, :, :]) + b[k, :, :]) local_weights = [np.zeros(rff_dim) for _ in range(num_clients)] # Initialize local weights for each client mse_values_per_iteration = np.zeros(num_iterations) # Store MSE for each iteration mse_values_per_iteration_per_client = np.zeros((num_clients, num_iterations)) # Store MSE for each client per iteration # Iterative training process for n in range(num_iterations): selected_indices = np.random.choice(num_clients, num_participating_clients, replace=False) # Select random clients for k in selected_indices: local_weights[k] = global_weights # Start with global weights epsilon = y[k, n, :] - np.dot(local_weights[k], z[k, n, :]) # Compute error local_weights[k] += learning_rate * z[k, n, :] * epsilon # Update local weights mse_values_per_iteration_per_client[k, n] = epsilon**2 # Compute MSE for the current iteration mse_values_per_iteration[n] += mse_values_per_iteration_per_client[k, n] # Aggregate MSE for selected clients mse_values_per_iteration[n] /= num_participating_clients # Average MSE over participating clients global_weights = np.zeros(rff_dim) # Reset global weights for k in selected_indices: global_weights += local_weights[k] # Aggregate updated local weights global_weights /= num_participating_clients # Average global weights mse_values_all_trials += mse_values_per_iteration # Accumulate MSE across all trials # Average MSE across all trials and normalize mse_values_all_trials /= independent_experiment mse_values_all_trials /= max(mse_values_all_trials) # Convert MSE to decibels mse_value_all_trials = 10 * np.log10(mse_values_all_trials) # Plot the MSE values over iterations plt.plot(mse_value_all_trials) plt.xlabel("Iterations") plt.ylabel("MSE (dB)") plt.title("Mean Squared Error Over Iterations") plt.show()
已尝试措施
- 严格按照论文方法实现基于RFF的KLMS联邦在线学习框架,覆盖客户端数据生成、本地更新与全局聚合全流程
- 验证AR模型数据生成逻辑及RFF特征变换的正确性
- 尝试调整学习率(0.75以外的大小值)以稳定MSE
- 确认全局模型聚合逻辑为选中客户端权重的平均
- 多次运行独立Monte Carlo试验消除随机性影响
预期结果
- MSE随迭代次数持续下降,体现全局模型的学习效果
- MSE曲线整体平滑,虽有波动但最终收敛至较低稳态值
- 结果与论文仿真图一致,收敛速率及稳态MSE匹配
内容的提问来源于stack exchange,提问作者Sunil Dhawan
相关产品推荐
相关产品推荐

