R语言中与Matlab mnrval函数等价的实现方法咨询
R语言等价实现Matlab
mnrval() 功能方案 Matlab中mnrval(B,X)的核心功能是基于多分类logistic回归的系数估计值,对输入的预测变量矩阵计算各分类的预测概率,核心逻辑不需要依赖复杂的包内部函数,可以直接通过原生R代码实现,也可以通过现有统计包函数适配。
核心逻辑说明
Matlab名义多分类logistic回归的mnrval计算规则:
- 输入
B为mnrfit返回的系数矩阵,维度为(P+1) × (K-1),其中P是预测变量个数,K是总分类数,默认以最后一个分类为参照类,第一行为截距项系数,后续P行对应每个预测变量的系数 - 输入
X为N × P的预测变量矩阵,函数自动添加截距列,不需要用户手动拼接全1列 - 输出为
N × K的预测概率矩阵,每行对应一个观测在K个类别上的概率,每行概率和为1
原生R实现代码(完全匹配基础功能)
以下实现不依赖任何R包的私有函数,输入输出规则和Matlab mnrval(B,X)完全对齐:
mnrval_r <- function(B, X) { # 统一X为矩阵格式 if (!is.matrix(X)) X <- as.matrix(X) n_obs <- nrow(X) p_var <- ncol(X) # 自动添加截距列,无需用户手动处理 X_aug <- cbind(1, X) # 维度校验 if (nrow(B) != p_var + 1) { stop("系数矩阵B维度不匹配,行数应为预测变量列数+1") } n_class <- ncol(B) + 1 # 计算线性预测值,参照类(最后一类)线性预测值固定为0 eta <- X_aug %*% B eta <- cbind(eta, 0) # softmax转换计算概率,做数值平移避免溢出 eta_shifted <- eta - apply(eta, 1, max) exp_eta <- exp(eta_shifted) phat <- exp_eta / rowSums(exp_eta) return(phat) }
使用说明
- 直接将Matlab
mnrfit导出的系数矩阵B、预测变量矩阵X传入mnrval_r(B,X)即可得到和Matlabmnrval完全一致的预测概率结果 - 注意将Matlab导出的矩阵在R中转换为标准
matrix类型,保持维度和Matlab中一致即可,不需要做转置或顺序调整 - 如果需要支持有序多分类(比例优势模型)的预测逻辑,只需要将上述代码中线性预测值、概率转换的部分替换为累积logit对应的计算规则即可,不需要依赖第三方包的内部接口。
内容的提问来源于stack exchange,提问作者A4747
相关产品推荐
相关产品推荐

