33 交互作用与非线性特征
33.1 交互作用
在开始建模过程之前,我们通常无法获知哪些预测变量存在交互作用。
我们可以通过可视化手段来探索潜在的交互作用。例如,使用散点图矩阵(scatterplot matrix)或条件散点图(conditional scatterplots)可以帮助我们识别变量之间的关系和潜在的交互作用。此外,还可以使用一些专门的工具量化、识别变量间的交互作用,其中Friedman 的 H 统计量是一个常用的方法,它是一类用于衡量模型中变量间交互作用强度的统计量,可以回答特征A和特征B在多大程度上不是”相互独立”的影响预测值这一问题。
H统计量越接近1,两个变量的交互作用越强;越接近0,交互作用越弱。
33.1.1 hstats包
hstats 是一个专门用来计算Friedman 的 H 统计量(及其变体)的 R 包。相比 iml,它速度更快、内存占用更低,且支持大型数据和并行计算,还能直接提供:
- 整体交互强度(各特征的总交互效应)。
- 两两交互强度(特征对之间的交互值)。
- 特征重要性(基于交互的分解)。
hstats()有三个核心输入:
-
预测函数或模型对象
object:支持lm, glm, gbm, ranger, xgboost, keras 等,或自定义函数。 -
背景数据
X:用于计算部分依赖的参考分布,通常为训练集或测试集。 -
目标变量
y:可选参数,用于计算特定类型的统计量。
33.1.1.1 计算H统计量
data("ames", package = "modeldata")
set.seed(123)
ames_sample <- ames |> slice_sample(n = 1000) # 随机抽样1000个样本
# 划分训练集和测试集
ames_split <- initial_split(ames_sample, prop = 0.8)
ames_train <- training(ames_split)
ames_test <- testing(ames_split)
# 训练随机森林模型
rf_fit <- rand_forest(mode = "regression") |>
set_engine("ranger") |>
fit(Sale_Price ~ ., data = ames_train)
# 提取预测函数
pred_fun <- \(object, newdata) {
predict(object, new_data = newdata) |> pull(.pred)
}
# 背景数据使用训练集
X <- ames_test |> select(-Sale_Price)
# 计算整体交互强度
ames_hstats_result <- hstats(
object = rf_fit,
X = X,
pred_fun = pred_fun,
n_max = 300, # 数据量过大时可选抽样
verbose = F
)33.1.1.2 提取和解读结果
# 提取整体交互强度:每个特征参与交互的总体水平
ames_hstats_result'hstats' object. Use plot() or summary() for details.
H^2 (normalized)
[1] 0.07520765
plot(ames_hstats_result) # 作图
# 提取结果数据框
overall_interaction <- ames_hstats_result |>
h2_overall(zero = TRUE)
# overall_interaction # 查看结果# 提取两两交互强度:特征对之间的交互值
pairwise_interaction <- ames_hstats_result |>
h2_pairwise()
pairwise_interactionPairwise H^2 (normalized)
[,1]
Mas_Vnr_Area:Garage_Area 0.0060172254
Gr_Liv_Area:Garage_Cars 0.0049202692
Mas_Vnr_Area:Garage_Cars 0.0045269802
Mas_Vnr_Area:Total_Bsmt_SF 0.0021544032
Gr_Liv_Area:Garage_Area 0.0019343259
Total_Bsmt_SF:Garage_Cars 0.0008562748
Total_Bsmt_SF:Gr_Liv_Area 0.0008160921
Garage_Cars:Garage_Area 0.0007365176
Total_Bsmt_SF:Garage_Area 0.0006554998
Mas_Vnr_Area:Gr_Liv_Area 0.0002371594
plot(pairwise_interaction)
- 约75.96%的预测变异无法有所有主效应的综合解释,说明模型中存在较强的交互作用。
- 最强的整体交互作用与
Garage_Area相关。 -
Mas_Vnr_Area:Total_Bsmt_SF是最强的两两交互作用,说明这两个特征之间存在显著的交互效应。
pd_importance(ames_hstats_result) |>
plot()
# 也可以提取数值
# importance_df <- ames_hstats_result |>
# perm_importance()-
交互值为负数:正常,\(H^2\) 可能因计算误差出现微小负值,可以
pmax(0, h2)修正或忽略。 -
特征数量很大:先用
h2_overall()筛选前k个重要的交互特征,再计算其两两交互,减少计算量。 -
分类变量:
hstats自动识别因子变量,但需要在背景数据中设定为因子类型,否则当作数值处理。 -
警告信息:如果出现 NaN 或超大值,通常因为某个特征方差极低或常数列,在建模前应剔除(
step_zv()很实用)。
33.2 基展开-Basis Expansion
基展开是将原始特征 ( X ) 映射到一组新的“基函数” ( \({b_1(X), b_2(X), \dots, b_d(X)}\)) 上,从而将模型从 ( f(X) ) 扩展为: \(f(X) = \sum_{m=1}^{d} \beta_m \cdot b_m(X)\)
原本的线性模型 ( \(y = \beta_0 + \beta_1 X\) ) 就是最简例(基函数为 \((1, X)\)),但基展开允许我们用非线性基函数来拟合曲线、曲面等复杂关系,而模型本身仍保持对系数 ( \(\beta\) ) 线性(即仍可用 OLS、正则化等线性方法求解)。
- 现实关系往往非线性:房价和面积的关系很少是直线,用多项式或样条才能更好拟合。
- 突破输入空间的限制:通过构造额外的“特征”,线性模型可以表达非线性函数(例如 ( X, X^2, X^3 ) 使得模型可拟合三次曲线)。
- 保持可解释性:尽管基函数可复杂,但最终模型仍是一个可加总形式,可逐一分析各基函数贡献(如上一节中可视化的样条基函数)。
- 结构化提升能力:能捕捉周期性、局域变化、交互项,而无需完全依赖非参数模型。
33.2.1 常见的基函数类型
| 类型 | 基函数形式 | 特点 | 适用场景 |
|---|---|---|---|
| 多项式 | (\(X,X^2,X^3,…\)) | 全局性强,高次容易边界震荡 | 低阶趋势拟合 |
| 分段多项式(回归样条) | 把X分段,每一段是低次多项式 | 灵活控制局部形状 | 连续数据的非线性关系 |
| B样条 | 由多个局部支撑的基函数构成(step_bs()) |
数值稳定,局部性,适合光滑拟合 | 通用样条建模 |
| 自然样条 | B样条+边界约束(边界外为线性) | 边界更稳定,减少外插风险(step_ns()) |
边界预测重要时 |
| 平滑样条 | 惩罚所有数据点+平滑参数 | 自动选择自由度 | 自动确定复杂度 |
| 径向基函数/核函数 | \(exp(−γ∥X−c∥2)\) | 局部拟合、可实现非线性映射 | 高维非线性回归、分类 |
| 傅里叶基 | 正弦、余弦序列 | 捕捉周期性 | 时间序列、季节分析 |
33.2.2 多项式基展开
# load data
data(deliveries, package = "modeldata")
set.seed(991)
delivery_split <- initial_validation_split(
deliveries,
prop = c(0.6, 0.2),
strata = time_to_delivery
)
delivery_train <- training(delivery_split)
delivery_test <- testing(delivery_split)
# step_poly()
recipe(time_to_delivery ~ ., data = delivery_train) |>
# 对数值列 hour 进行 4 次正交多项式展开(默认生成 hour_poly_1、hour_poly_2、hour_poly_3、hour_poly_4 共 4 个新列),
# 并移除原始 hour 列(因为 keep_original_cols 默认为 FALSE)
step_poly(hour, degree = 4) |>
prep() |>
# 将训练好的预处理应用到数据上。new_data = NULL 表示处理 prep() 时所用的原始数据(delivery_train)
bake(new_data = NULL, starts_with("hour"))# A tibble: 6,004 × 4
hour_poly_1 hour_poly_2 hour_poly_3 hour_poly_4
<dbl> <dbl> <dbl> <dbl>
1 -0.0232 0.0224 -0.0176 0.00800
2 0.0156 0.0103 -0.00312 -0.0134
3 0.0111 -0.000374 -0.0117 -0.0104
4 -0.00237 -0.0132 0.00482 0.0118
5 -0.0176 0.00628 0.00936 -0.0173
6 -0.0232 0.0225 -0.0178 0.00825
7 0.000437 -0.0128 -0.00103 0.0118
8 0.00821 -0.00542 -0.0121 -0.00348
9 -0.00262 -0.0132 0.00536 0.0116
10 0.0163 0.0121 -0.000827 -0.0123
# ℹ 5,994 more rows
# step_poly_bernstein()
recipe(time_to_delivery ~ ., data = delivery_train) |>
step_poly_bernstein(hour, degree = 4) |>
prep() |>
bake(new_data = NULL, starts_with("hour"))# A tibble: 6,004 × 4
hour_1 hour_2 hour_3 hour_4
<dbl> <dbl> <dbl> <dbl>
1 0.259 0.0359 0.00221 0.0000509
2 0.0167 0.121 0.390 0.472
3 0.0511 0.220 0.422 0.303
4 0.266 0.374 0.234 0.0549
5 0.405 0.144 0.0227 0.00134
6 0.258 0.0355 0.00217 0.0000497
7 0.212 0.371 0.288 0.0838
8 0.0841 0.277 0.406 0.223
9 0.271 0.374 0.229 0.0527
10 0.0134 0.107 0.378 0.502
# ℹ 5,994 more rows
- 如果需要对不同列使用不同的多项式次数,则需多次调用该步骤。
33.2.3 样条函数
apropos("step_spline")[1] "step_spline_b" "step_spline_convex"
[3] "step_spline_monotone" "step_spline_natural"
[5] "step_spline_nonnegative"
- 每个样条函数的语法与步骤几乎都是相同的。
- 样条函数极为实用,尤其当我们希望促使简单模型(如线性回归)逼近更复杂的黑箱模型(例如神经网络或树集成)的预测性能时,其便利性尤为突出。
recipe(time_to_delivery ~ ., data = delivery_train) |>
# 三次样条是常见选择(deg_free),因为相比线性或二次拟合,它提供了更大的灵活性,但又不会过于灵活而导致过拟合。
step_spline_natural(hour, deg_free = 4) |>
prep() |>
bake(new_data = NULL, starts_with("hour"))# A tibble: 6,004 × 4
hour_1 hour_2 hour_3 hour_4
<dbl> <dbl> <dbl> <dbl>
1 0.197 0.00430 0 0
2 0 0.0633 0.431 0.312
3 0 0.217 0.482 0.261
4 0.203 0.666 0.0362 0.0167
5 0.394 0.0500 0 0
6 0.196 0.00422 0 0
7 0.103 0.723 0.0971 0.0448
8 0.000850 0.381 0.417 0.206
9 0.213 0.656 0.0324 0.0149
10 0 0.0500 0.410 0.317
# ℹ 5,994 more rows
- 方差-偏差权衡是机器学习中理解模型预测误差来源的核心概念,它揭示了一个模型在欠拟合与过拟合之间必须做出的取舍。
- 偏差(Bias):平均预测值与真实函数的偏离程度。反映了模型对数据背后规律的系统性简化。高偏差意味着模型过于简单,无法捕捉真实模式(欠拟合)。
- 方差(Variance):在不同训练集上训练出的模型,在 ( \(x_0\) ) 处预测值的波动程度。反映了模型对训练数据的敏感度。高方差意味着模型过度追踪训练集中的噪声和随机波动(过拟合)。
- 不可约误差(Irreducible Error):由数据本身噪声引起,无法通过任何模型消除。
在tidymodels中:
- 使用 recipes 进行基展开(
step_poly,step_ns)就是在降低偏差。 - 使用 glmnet 或 parsnip 中的正则化模型就是在控制方差。
- 使用 rsample + tune 进行交叉验证,就是在寻找最优的偏差-方差平衡点。
- 下面两个图展示了方差-偏差发生变化时预测值产生的分布变化。章节 27 的开头部分也对这个概念进行了说明。
,
33.3 离散化
离散化是将连续型数值变量转换为分类变量(因子)的过程。核心思想是把数值范围切分成有限的几个区间,每个区间对应一个类别。离散化在特征工程中有多种用途:
- 处理非线性:将连续变量分段后,可用简单的线性模型捕捉分段关系(类似阶梯函数)。
- 增强鲁棒性:减少异常值和噪声的影响(噪声数据被归并到区间内)。
- 改善可解释性:例如把年龄变为“少年、青年、中年、老年”比原始年龄更易理解。
- 满足算法要求:部分算法(如朴素贝叶斯、关联规则)只接受分类数据。
- 但离散化也会丢失信息(区间内部差异被忽略),所以需谨慎使用,通常应交叉验证确定最优分箱数。
在recipe中:
- 使用
step_discretize()进行无监督分箱1。 - 使用
embed::step_discretize_cart()和embed::step_discretize_xbg()进行有监督分箱2。 - 如果希望使用自定义断点,可使用
recipe::step_cut()。
basic_binning <-
recipe(time_to_delivery ~ ., data = delivery_train) |>
step_discretize(hour, num_breaks = 4) |>
prep()
bake(basic_binning, new_data = NULL) |>
select(hour)# A tibble: 6,004 × 1
hour
<fct>
1 bin1
2 bin4
3 bin4
4 bin2
5 bin1
6 bin1
7 bin2
8 bin3
9 bin2
10 bin4
# ℹ 5,994 more rows
# 检查分箱断点位置
tidy(basic_binning, 1)# A tibble: 5 × 3
terms value id
<chr> <dbl> <chr>
1 hour -Inf discretize_viGXH
2 hour 14.5 discretize_viGXH
3 hour 16.6 discretize_viGXH
4 hour 18.2 discretize_viGXH
5 hour Inf discretize_viGXH
# A tibble: 6,004 × 1
hour
<fct>
1 [-Inf,12.28)
2 [18.97,19.68)
3 [16.63,18.97)
4 [15.73,16.24)
5 [12.28,13.23)
6 [-Inf,12.28)
7 [16.24,16.63)
8 [16.63,18.97)
9 [15.73,16.24)
10 [18.97,19.68)
# ℹ 5,994 more rows
# 检查分箱断点位置
tidy(cart_binning, 1)# A tibble: 11 × 3
terms value id
<chr> <dbl> <chr>
1 hour 12.3 discretize_cart_xzO5P
2 hour 13.2 discretize_cart_xzO5P
3 hour 14.0 discretize_cart_xzO5P
4 hour 14.8 discretize_cart_xzO5P
5 hour 15.7 discretize_cart_xzO5P
6 hour 16.2 discretize_cart_xzO5P
7 hour 16.6 discretize_cart_xzO5P
8 hour 19.0 discretize_cart_xzO5P
9 hour 19.7 discretize_cart_xzO5P
10 hour 19.8 discretize_cart_xzO5P
11 hour 20.3 discretize_cart_xzO5P