使用partykit的MOB过程预测生存概率遇报错及参数疑问
问题与解决方案
一、预测错误的修复
你遇到的错误确实是测试集与训练集各终端节点观测数不匹配导致的,predict.mob在处理分位数预测时无法自动对齐维度。可通过手动分配节点+逐个节点预测的方式解决:
解决步骤
- 确定测试集每个观测所属的终端节点
- 提取MOB树中每个终端节点的
survreg模型 - 对每个节点内的测试集观测,调用对应模型的分位数预测
- 按原始顺序整理预测结果
修改后的预测代码片段
# 确定测试集观测的节点归属 node_ids <- predict(glmtr, newdata = dat.test, type = "node") # 提取所有终端节点的模型 terminal_nodes <- glmtr$node[glmtr$node$terminal] models <- lapply(terminal_nodes, function(n) n$fit) names(models) <- names(terminal_nodes) # 初始化预测结果矩阵 quantile_pred <- matrix(NA, nrow = nrow(dat.test), ncol = length(pct)) # 逐个节点预测并填充结果 for(node in unique(node_ids)) { idx <- which(node_ids == node) # 调用对应节点的survreg模型预测分位数 pred <- predict(models[[as.character(node)]], newdata = dat.test[idx, ], type = "quantile", p = pct) quantile_pred[idx, ] <- pred }
二、parm选项的正确性确认
你当前设置parm=2是错误的,原因如下:
- 你的
wbreg函数使用survreg(y ~ 0 + x),其中x对应模型公式Surv(Y, D) ~ W + X1 + X2的设计矩阵,列顺序为W、X1、X2 parm参数指定需要检验参数不稳定性的系数索引,要检验处理效应W的参数不稳定性,应设置parm=1
修改后的mob模型拟合代码:
glmtr <- partykit::mob(as.formula(eqn), data = dat, fit = wbreg, control = mob_control(parm=1, minsize = 0.2*nrow(dat), alpha = 0.10, bonferroni = TRUE))
完整修改后代码
library("survival") library("partykit") n=5000;n.test=5000;p=25;pi=0.5;beta=1 gamma=0.5;rho=2;cen.scale=4;n.mc=10000; Y.max=2 generate_data <- function(n, p, pi = 0.5, beta = 1, gamma = 1, rho = 2, cen.scale = 4, Y.max = NULL){ W <- rbinom(n, 1, pi) X <- matrix(rnorm(n * p), n, p) numerator <- -log(runif(n)) cox.ft <- (numerator / exp(beta * X[ ,1] + (-0.5 - gamma * X[ ,2]) * W))^2 failure.time <- pmin(cox.ft, Y.max) numeratorC <- -log(runif(n)) censor.time <- (numeratorC / (cen.scale ^ rho)) ^ (1 / rho) Y <- pmin(failure.time, censor.time) D <- as.integer(failure.time <= censor.time) list(X = X, Y = Y, W = W, D = D) } data <- generate_data(n, p=p, pi = pi, beta = beta, gamma = gamma, rho = rho, cen.scale = cen.scale, Y.max = Y.max) data.test <- generate_data(n.test, p=p, pi = pi, beta = beta, gamma = gamma, rho = rho, cen.scale = cen.scale, Y.max = Y.max) X=data$X Y=data$Y W=data$W D=data$D var_prog <- c("X1","X2") colnames(X) <- paste("X", 1:25, sep="") cov.names <- colnames(X) wbreg <- function(y, x, start = NULL, weights = NULL, offset = NULL, ...) { survreg(y ~ 0 + x, weights = weights, dist = "weibull", ...) } dat <- data.frame(Y=Y,D=D,W=W,X) eqn <- paste0("Surv(Y, D) ~ W + ",paste0(var_prog, collapse="+")," | ", paste0(cov.names, collapse="+")) # 修正parm参数为1,对应W的系数 glmtr <- partykit::mob(as.formula(eqn), data = dat, fit = wbreg, control = mob_control(parm=1, minsize = 0.2*nrow(dat), alpha = 0.10, bonferroni = TRUE)) plot(glmtr) dat.test <- data.frame(Y=data.test$Y,D=data.test$D, W=data.test$W,data.test$X) colnames(dat.test) <- colnames(dat) # 确保测试集列名与训练集完全一致 pct <- 1:98/100 # 修正后的预测流程 node_ids <- predict(glmtr, newdata = dat.test, type = "node") terminal_nodes <- glmtr$node[glmtr$node$terminal] models <- lapply(terminal_nodes, function(n) n$fit) names(models) <- names(terminal_nodes) quantile_pred <- matrix(NA, nrow = nrow(dat.test), ncol = length(pct)) for(node in unique(node_ids)) { idx <- which(node_ids == node) pred <- predict(models[[as.character(node)]], newdata = dat.test[idx, ], type = "quantile", p = pct) quantile_pred[idx, ] <- pred }
内容的提问来源于stack exchange,提问作者Remy M
相关产品推荐
相关产品推荐

