You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何扩展nnet::multinom的predict方法以支持type="link"适配marginaleffects

问题背景

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)
}

关键修改说明

  1. 参数扩展:在type参数中新增"link"选项,确保match.arg能识别
  2. newdata存在时的逻辑:
    • 调用predict.nnet时指定type="raw",拿到未经过softmax激活的线性预测值
    • type="probs"时对线性预测值应用softmax转换为概率,替代原方法直接调用predict.nnet返回概率的逻辑
    • type="link"时直接返回线性预测值
  3. 无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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.25 23:18:24