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

如何获取StatsModels.ModelMatrix变量名并复现数据映射适配分类值差异

问题描述

现有如下Julia DataFrame数据,需基于指定公式执行回归分析:

using GLM, StatsModels, Tables, DataFrames
training = DataFrame(yy = [1,2,3,7,5,4,2,3], continuous = [5,5,6,6,7,8,8,9], categorical = [:a,:a,:b,:b, :a,:c,:c,:b], bool = [true, false, true, true, true, false, false, true])
f = Term(:yy)~Term(:continuous) + Term(:categorical) + Term(:bool)

通过以下代码手动构建特征矩阵(因实际使用的回归模型无直接传入公式和DataFrame的接口):

cols = Tables.columntable(training)
mf = StatsModels.ModelFrame(f, cols, model=GLM.LinearModel)
mm = StatsModels.ModelMatrix(mf)
fitted = fit(GLM.LinearModel, mm.m, response(mf))

但后续用拟合模型预测时,测试集可能包含与训练集不同的分类值(如下方测试集不含:b),直接用StatsModels.ModelMatrix构建的矩阵会与训练集矩阵维度/特征不匹配:

test = DataFrame(continuous = [5,5,6,6,7,8,8,9], categorical = [:a,:a,:a,:a, :a,:c,:c,:a], bool = [false, false, false, false, false, false, false, false])

需要解决两个问题:

  1. 获取StatsModels.ModelMatrix.m矩阵的变量名
  2. 复现从DataFrame到矩阵的映射逻辑,适配与训练集分类值不同的测试集
解决方案

1. 获取ModelMatrix的变量名

ModelFrame对象的schema属性包含特征编码信息,结合ModelMatrix的赋值规则可提取每一列对应的变量名:

# 初始化特征名数组
feature_names = String[]

# 遍历训练集ModelFrame的所有项
for term in mf.schema[].terms
    # 处理分类变量的虚拟编码列
    if term isa StatsModels.CategoricalTerm
        levels = StatsModels.levels(term)
        # 跳过参考水平(默认是第一个分类值)
        for level in levels[2:end]
            push!(feature_names, "$(term.sym)_$(level)")
        end
    else
        # 连续变量/布尔变量直接取变量名
        push!(feature_names, String(term.sym))
    end
end

# 如果模型包含截距项,添加到最前面
if hasintercept(mf)
    pushfirst!(feature_names, "(Intercept)")
end

执行后feature_names即为mm.m矩阵每一列对应的变量名,示例中会得到:
["(Intercept)", "continuous", "categorical_b", "categorical_c", "bool"]

2. 适配测试集的特征映射逻辑

核心是复用训练集ModelFrame的schema,强制测试集遵循相同的编码规则,确保特征矩阵结构一致:

# 保存训练集的schema(关键:记录训练时的编码规则)
train_schema = mf.schema

# 用训练集schema处理测试集
test_cols = Tables.columntable(test)
test_mf = StatsModels.ModelFrame(f, test_cols; schema=train_schema, model=GLM.LinearModel)
test_mm = StatsModels.ModelMatrix(test_mf)

生成的test_mm.m矩阵维度和特征顺序会与训练集完全匹配:如果测试集缺少训练集的分类水平,对应列会填充0;若测试集出现训练集没有的分类值,代码会直接报错(可提前过滤测试集的异常分类值)。

比如示例中的测试集不含:b,test_mm.m里的categorical_b列会全为0,完全适配训练集的特征矩阵结构。

内容的提问来源于stack exchange,提问作者Stuart

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 09:27:45