如何获取R中Dirichlet回归预测值的标准误?
解决DirichletReg中预测值标准误的计算问题
针对DirichletReg包(v0.7-1)predict方法无法输出标准误的问题,我们可以通过delta方法(利用参数协方差矩阵推导非线性函数的方差)手动计算,下面分两种参数化方式逐一说明:
一、Common参数化(每个类别单独建模)
模型背景
Common参数化下,每个类别j的α参数满足log(α_j) = Xβ_j,最终预测均值为:μ_j = exp(Xβ_j) / sum_{k=1}^C exp(Xβ_k)
精度φ则是所有α的和:φ = sum(exp(Xβ_j))
计算步骤
核心思路是对每个类别预测均值求参数的偏导,再结合协方差矩阵计算方差:
- 提取模型的参数估计和协方差矩阵;
- 生成新数据对应的模型矩阵;
- 计算每个类别预测均值对所有参数的梯度向量;
- 用梯度向量和协方差矩阵计算方差,开平方得到标准误。
代码实现
# 拟合common参数化模型 model_common <- DirichletReg(y ~ x1 + x2, data = your_data, model = "common") # 准备新数据 new_dat <- data.frame(x1 = c(1, 2), x2 = c(3, 4)) # 示例新数据 # 生成匹配模型结构的预测矩阵(自动处理截距、因子编码) X_new <- model.matrix(delete.response(terms(model_common)), new_dat) # 提取关键模型输出 coefs <- coef(model_common) V <- vcov(model_common) C <- length(model_common$varnames$response) # 获取类别总数 beta_names <- grep("beta", names(coefs), value = TRUE) # 批量计算每个样本、每个类别的预测值和标准误 results <- list(pred_mu = matrix(nrow = nrow(new_dat), ncol = C), se_mu = matrix(nrow = nrow(new_dat), ncol = C)) for (i in 1:nrow(new_dat)) { x_i <- X_new[i,, drop = FALSE] # 计算所有类别的alpha线性预测值 eta_all <- sapply(1:C, function(j) { x_i %*% coefs[grepl(paste0("beta", j), beta_names)] }) alpha_all <- exp(eta_all) mu_i <- alpha_all / sum(alpha_all) results$pred_mu[i,] <- mu_i # 计算每个类别的梯度向量 for (j in 1:C) { grad <- numeric(length(coefs)) for (k in 1:C) { k_idx <- grepl(paste0("beta", k), names(coefs)) if (k == j) { grad[k_idx] <- mu_i[j] * (1 - mu_i[j]) * x_i } else { grad[k_idx] <- -mu_i[j] * mu_i[k] * x_i } } # 计算标准误 results$se_mu[i,j] <- sqrt(t(grad) %*% V %*% grad) } } # 查看结果 print(results$pred_mu) print(results$se_mu)
二、Alternative参数化(C-1类别建模+精度单独建模)
模型背景
Alternative参数化下,均值部分采用多项logit转换:
- 前C-1个类别:
log(μ_j / μ_C) = Xβ_j(μ_C为参考类别) - 预测均值:
μ_j = exp(Xβ_j)/(1 + sum_{k=1}^{C-1} exp(Xβ_k)),μ_C = 1/(1 + sum_{k=1}^{C-1} exp(Xβ_k))
精度部分单独建模:log(φ) = Zγ
计算步骤
均值的标准误计算同样基于delta方法,注意精度参数不影响均值,仅需对β参数求偏导;如果需要精度的标准误,单独处理即可。
代码实现
# 拟合alternative参数化模型(|后是精度模型的公式) model_alt <- DirichletReg(y ~ x1 + x2 | x1 + x2, data = your_data, model = "alternative") # 准备新数据 new_dat <- data.frame(x1 = c(1, 2), x2 = c(3, 4)) # 分别生成均值和精度部分的模型矩阵 X_new <- model.matrix(delete.response(terms(model_alt$formula$mean)), new_dat) Z_new <- model.matrix(delete.response(terms(model_alt$formula$precision)), new_dat) # 提取关键模型输出 coefs <- coef(model_alt) V <- vcov(model_alt) C <- length(model_alt$varnames$response) beta_names <- grep("beta", names(coefs), value = TRUE) gamma_names <- grep("gamma", names(coefs), value = TRUE) # 批量计算预测值和标准误 results <- list(pred_mu = matrix(nrow = nrow(new_dat), ncol = C), se_mu = matrix(nrow = nrow(new_dat), ncol = C), pred_phi = numeric(nrow(new_dat)), se_phi = numeric(nrow(new_dat))) for (i in 1:nrow(new_dat)) { x_i <- X_new[i,, drop = FALSE] z_i <- Z_new[i,, drop = FALSE] # 计算均值预测值 eta_beta <- sapply(1:(C-1), function(j) { x_i %*% coefs[grepl(paste0("beta", j), beta_names)] }) exp_eta <- exp(eta_beta) sum_exp <- sum(exp_eta) mu_i <- c(exp_eta/(1 + sum_exp), 1/(1 + sum_exp)) results$pred_mu[i,] <- mu_i # 计算均值的标准误 # 处理前C-1个类别 for (j in 1:(C-1)) { grad <- numeric(length(coefs)) for (k in 1:(C-1)) { k_idx <- grepl(paste0("beta", k), names(coefs)) if (k == j) { grad[k_idx] <- mu_i[j] * (1 - mu_i[j]) * x_i } else { grad[k_idx] <- -mu_i[j] * mu_i[k] * x_i } } # 精度参数的偏导为0,无需处理 results$se_mu[i,j] <- sqrt(t(grad) %*% V %*% grad) } # 处理参考类别 grad_C <- numeric(length(coefs)) for (k in 1:(C-1)) { k_idx <- grepl(paste0("beta", k), names(coefs)) grad_C[k_idx] <- -mu_i[C] * mu_i[k] * x_i } results$se_mu[i,C] <- sqrt(t(grad_C) %*% V %*% grad_C) # 计算精度的预测值和标准误(可选) eta_gamma <- z_i %*% coefs[gamma_names] phi_i <- exp(eta_gamma) results$pred_phi[i] <- phi_i # 精度的梯度向量 grad_phi <- numeric(length(coefs)) grad_phi[grepl("gamma", names(coefs))] <- phi_i * z_i results$se_phi[i] <- sqrt(t(grad_phi) %*% V %*% grad_phi) } # 查看结果 print(results$pred_mu) print(results$se_mu) print(results$pred_phi) print(results$se_phi)
重要注意事项
- 必须用
model.matrix()基于模型公式生成新数据的预测矩阵,确保和拟合模型时的变量编码(比如因子、截距)完全一致,避免手动构建出错; vcov(model)返回的是ML估计的参数协方差矩阵,已经考虑了参数间的相关性,delta方法的计算是可靠的;- 如果模型使用了稳健方差估计,需要替换为对应的协方差矩阵(DirichletReg默认是ML协方差)。
内容的提问来源于stack exchange,提问作者M. Riera
相关产品推荐
相关产品推荐

