如何用经验似然实现最大似然估计?求PyTorch/Jax实现示例
经验似然在无参数化似然回归问题中的PyTorch实现示例
核心思路
在无需指定参数化似然的回归任务中,经验似然的核心是为每个训练样本分配非负权重(权重和为1),通过最大化加权经验似然来拟合模型,同时满足回归残差的约束。嵌套坐标下降法将优化拆分为两层:
- 内层:固定回归参数,用坐标下降更新样本权重,满足权重约束与残差的经验似然条件
- 外层:固定权重,用梯度下降更新回归参数
PyTorch实现代码
import torch import numpy as np def empirical_likelihood_coordinate_descent(X, y, max_iter=100, tol=1e-6): # 初始化线性回归参数:y = X@beta beta = torch.randn(X.shape[1], 1, requires_grad=True) optimizer = torch.optim.SGD([beta], lr=0.01) # 初始化样本权重为均匀分布 weights = torch.ones(X.shape[0], 1) / X.shape[0] weights.requires_grad = False prev_loss = float('inf') for outer_iter in range(max_iter): # 外层:固定权重,优化回归参数 optimizer.zero_grad() y_pred = X @ beta # 负对数加权经验似然损失(无需指定残差分布) loss = -torch.sum(weights * torch.log(torch.clamp(1e-8 + torch.abs(y - y_pred), min=1e-8))) loss.backward() optimizer.step() # 内层:固定beta,坐标下降更新权重 with torch.no_grad(): y_pred_fixed = X @ beta residuals = y - y_pred_fixed for inner_iter in range(50): old_weights = weights.clone() for i in range(X.shape[0]): sum_rest = torch.sum(weights) - weights[i] # 基于经验似然KKT条件推导的权重更新公式 numerator = 1.0 / torch.abs(residuals[i]) denominator = torch.sum(1.0 / torch.abs(residuals)) - numerator weights[i] = numerator / denominator * (1 - sum_rest) weights[i] = torch.clamp(weights[i], min=1e-8) # 内层迭代收敛判断 if torch.norm(weights - old_weights) < tol: break # 外层迭代收敛判断 if torch.abs(loss.item() - prev_loss) < tol: break prev_loss = loss.item() return beta.detach(), weights.detach() # 测试用例 if __name__ == "__main__": # 生成模拟回归数据:y = 2x1 + 3x2 + 高斯噪声 np.random.seed(42) X_np = np.random.randn(100, 2) y_np = 2*X_np[:,0] + 3*X_np[:,1] + np.random.randn(100)*0.5 X = torch.tensor(X_np, dtype=torch.float32) y = torch.tensor(y_np.reshape(-1,1), dtype=torch.float32) beta_est, weights_est = empirical_likelihood_coordinate_descent(X, y) print("估计的回归系数:", beta_est.numpy().flatten()) print("前5个样本权重:", weights_est.numpy().flatten()[:5])
关键部分解释
- 权重约束处理:坐标下降过程中强制权重非负且总和为1,避免出现无效的负权重或权重和失衡
- 嵌套优化逻辑:外层利用PyTorch自动微分处理回归参数的梯度更新,内层用无梯度循环执行坐标下降,完全匹配论文中的嵌套策略
- 无参数似然设计:损失函数基于残差绝对值的对数加权和,无需假设残差服从特定参数化分布,符合无参数化似然的回归场景要求
Jax实现思路
Jax版本可借助其函数式编程和自动向量化特性简化逻辑:
- 用
jax.lax.fori_loop实现内层坐标下降的高效循环 - 用
jax.grad计算回归参数的梯度 - 权重更新逻辑与PyTorch版本一致,仅需替换为Jax的张量操作
内容的提问来源于stack exchange,提问作者Ggjj11
相关产品推荐
相关产品推荐

