如何扩展nnet::multinom的predict方法以支持type="link"适配marginaleffects
让
nnet::multinom兼容marginaleffects:扩展predict.multinom支持type="link" 问题背景
marginaleffects::predictions()依赖模型的predict()方法同时提供响应尺度和链接尺度的预测值,但nnet包原生的predict.multinom仅支持type="probs"或type="class"参数,缺少链接尺度(type="link")的输出能力,导致无法与marginaleffects兼容。
实现方案
我们可以通过重定义predict.multinom方法,新增对type="link"的支持,核心思路是:
- 扩展
type参数选项,加入"link" - 复用原方法的数据处理逻辑,在结果转换阶段新增链接尺度的计算:调用
predict.nnet获取softmax激活前的线性预测值,直接作为链接尺度输出;对概率和类别输出逻辑做对应调整以保持一致性
修改后的predict.multinom代码
predict.multinom <- function(object, newdata, type=c("class","probs", "link"), ...) { if(!inherits(object, "multinom")) stop("not a \"multinom\" fit") type <- match.arg(type) if(missing(newdata)) { Y <- fitted(object) # 从拟合概率反推链接尺度预测值(以第一个类别为参照) if(type == "link") { Y <- log(Y[, -1, drop=FALSE] / Y[, 1, drop=FALSE]) Y <- cbind(0, Y) colnames(Y) <- object$lev } } else { newdata <- as.data.frame(newdata) rn <- row.names(newdata) Terms <- delete.response(object$terms) m <- model.frame(Terms, newdata, na.action = na.omit, xlev = object$xlevels) if (!is.null(cl <- attr(Terms, "dataClasses"))) .checkMFClasses(cl, m) keep <- match(row.names(m), rn) X <- model.matrix(Terms, m, contrasts = object$contrasts) # 获取softmax激活前的线性预测值 Y1 <- predict.nnet(object, X, type = "raw") Y <- matrix(NA, nrow(newdata), ncol(Y1), dimnames = list(rn, object$lev)) Y[keep, ] <- Y1 # 根据type参数转换输出 switch(type, probs = { Y <- exp(Y) / rowSums(exp(Y)) }, class = { prob <- exp(Y) / rowSums(exp(Y)) if(length(object$lev) > 2L) Y <- factor(max.col(prob), levels=seq_along(object$lev), labels=object$lev) if(length(object$lev) == 2L) Y <- factor(1 + (prob > 0.5), levels=1L:2L, labels=object$lev) if(length(object$lev) == 0L) Y <- factor(max.col(prob), levels=seq_along(object$lab), labels=object$lab) }, link = {} # 直接保留线性预测值 ) } drop(Y) }
关键修改说明
- 参数扩展:在
type参数中新增"link"选项,确保match.arg能识别 newdata存在时的逻辑:- 调用
predict.nnet时指定type="raw",拿到未经过softmax激活的线性预测值 type="probs"时对线性预测值应用softmax转换为概率,替代原方法直接调用predict.nnet返回概率的逻辑type="link"时直接返回线性预测值
- 调用
- 无
newdata时的逻辑:从拟合的概率值反推链接尺度的log-odds(以第一个类别为参照),并补回参照类别的0值
验证示例
# 加载数据 dat = read.csv("https://www.dropbox.com/s/u27cn44p5srievq/dat.csv?dl=1") dat$collection_date = as.Date(dat$collection_date) dat$collection_date_num = as.numeric(dat$collection_date) dat$variant = factor(dat$variant) # 加载依赖包 library(nnet) library(splines) library(marginaleffects) # 拟合模型 set.seed(1) fit_nnet = nnet::multinom(variant ~ ns(collection_date_num, df=2), weights=count, data=dat) # 测试链接尺度预测 link_preds <- predict(fit_nnet, newdata = datagrid(collection_date_num = max(dat$collection_date_num)), type="link") print(link_preds) # 使用marginaleffects计算预测概率及置信区间 multinom_preds_marginaleffects = predictions(fit_nnet, newdata = datagrid(collection_date_num = max(dat$collection_date_num)), type="link", transform_post = insight::link_inverse(fit_nnet)) print(multinom_preds_marginaleffects)
内容的提问来源于stack exchange,提问作者Tom Wenseleers
相关产品推荐
相关产品推荐

