如何用 LightGBM 同时训练多个相关目标:多任务学习实战指南
【免费下载链接】LightGBMA fast, distributed, high performance gradient boosting (GBT, GBDT, GBRT, GBM or MART) framework based on decision tree algorithms, used for ranking, classification and many other machine learning tasks.项目地址: https://gitcode.com/GitHub_Trending/li/LightGBM
电网调度会前,要的不是一个数,而是三个数:次日平均负荷、峰值负荷、谷值负荷。用 LightGBM 做多任务学习,不必为这三个目标各训一个模型,一次训练就能同时拿到三份预测。
为什么值得一次做多个目标
最常见的做法是一个目标训一个模型。每个模型都没问题,但 k 个目标意味着 k 轮训练、k 个模型文件、k 套监控。更关键的是,如果这些目标来自同一套特征,模型会把共同结构重复学 k 遍,目标之间的信息一点都没用上。
💡 可以这样理解多任务学习:像同部门的几个新人各管各的项目,但共用一个资料库和流程文档,公共积累不用从零重来。收益大小取决于任务相关性——两个目标的波动越同步,一起训练的收益越大;目标之间不相关,强行合并只会增加复杂度。
单任务和多任务方案的差异,直接看表:
| 维度 | 每个目标独立建模 | 多任务联合训练 |
|---|---|---|
| 模型数量与训练成本 | k 个模型、k 轮训练 | 1 个模型(或轻量包装),1 轮训练 |
| 任务间信息利用 | 无,各自独立学习 | 共享树结构与特征分裂,相关目标互相借力 |
| 预测一致性 | 各任务结果可能互相矛盾 | 联合优化,天然保持任务间协调 |
| 工程维护 | 模型产物、部署、监控都是 k 份 | 单一模型产物,统一迭代 |
| 适用前提 | 任意任务组合 | 目标同类型、特征共享、且相关性足够 |
LightGBM 多任务学习实现路径
LightGBM 没有"多任务模式"开关,但用官方 Python 包就能搭出三条从轻到重的路径,按任务结构选一条即可。
最轻:用 sklearn 的 MultiOutput 按任务包装
sklearn.multioutput里的MultiOutputClassifier/MultiOutputRegressor会把你给的一个 LightGBM 估计量按任务克隆、独立训练。好处是二分类和回归可以混在同一个流程里,训练代码几乎不用改。
import lightgbm as lgb from sklearn.multioutput import MultiOutputRegressor from sklearn.model_selection import train_test_split Xtr, Xte, Ytr, Yte = train_test_split(X, Y, random_state=42) model = MultiOutputRegressor(lgb.LGBMRegressor(n_estimators=300, verbose=-1)) model.fit(Xtr, Ytr) # 按 Y 的列数逐个训练 LGBMRegressor pred = model.predict(Xte) # shape (n_test, k),每列对应一个任务这套方案里 k 个模型互相独立,可以丢给并行任务(后文性能调优会讲);所谓"共享"的是数据管道和评估流程,而不是模型结构本身。
中间档:一个模型多目标,自定义目标函数加权
如果所有任务都是二分类,LightGBM 有内置多标签目标:把 y 传成 n×k 的 0/1 矩阵,配objective="multilabel_binary",predict 直接返回每个任务的概率——这是最标准的"一个模型、多个二分类任务"写法。
对相关的回归任务,可以用标签堆叠技巧:把每个样本的特征纵向重复 k 次,k 个目标堆成一列,一个回归模型同时拟合所有任务。此时把 objective 传成 callable,就能写多任务损失,比如给主任务更高权重:
k = 3 X_stack = np.tile(Xtr, (k, 1)) # 特征纵向重复 k 份 y_stack = Ytr.T.reshape(-1) # 按任务堆叠:先任务1全体,再任务2… def multi_task_objective(y_true, y_pred): n = len(y_true) // k task_id = np.repeat(np.arange(k), n) # 每行属于哪个任务 w = np.array([1.0, 2.0, 1.0])[task_id] # 主任务权重更高 grad = 2.0 * w * (y_pred - y_true) hess = 2.0 * w return grad, hess params = {"objective": multi_task_objective, "metric": "l2", "verbose": -1} bst = lgb.train(params, lgb.Dataset(X_stack, y_stack), num_boost_round=300) pred = bst.predict(np.tile(Xte, (k, 1))).reshape(k, -1).T # (n, k)这样可行的原因:树从同一套特征上学分裂,分裂点由所有任务的加权梯度共同决定,特征里的公共模式一次分裂就惠及全部任务。自定义目标函数正是调节"每个任务投入多少"的入口,加权、换损失函数都能在这里改。
最重:任务标识特征,一个模型覆盖多任务
标签堆叠是"同一份样本出多个目标";如果不同任务的特征集不一样,或者任务很多,更灵活的是"样本展开 + 任务特征":把所有任务的样本拼成一张表,加一列task_id当特征。树靠这个特征做特化,其余特征的分裂仍然复用。
n = X.shape[0] X_feat = np.column_stack([np.vstack([X, X]), np.concatenate([np.zeros(n), np.ones(n)])]) y_feat = np.concatenate([Y[:, 0], Y[:, 1]]) # 这里演示 2 个任务 model = lgb.LGBMRegressor(n_estimators=300, verbose=-1) model.fit(X_feat, y_feat) pred_0 = model.predict(np.column_stack([Xte, 0])) # 任务 0 的预测 pred_1 = model.predict(np.column_stack([Xte, 1])) # 任务 1 的预测代价是数据量翻倍、要维护特征对齐;换来的是单一产物,新增任务只需多一个task_id,不用改模型结构。
怎么判断多任务模型有没有用
值不值得合并,看两样证据。一是任务相关性:算目标间的相关系数矩阵,非对角值越高(经验上 |r| > 0.5 是参考线),联合训练的预期收益越大。二是逐任务评估:只看平均指标会掩盖单任务劣化,要按任务报指标,并和独立模型逐一对照——多任务模型只有在每个任务上都不输基线、且优势超出噪声,才值得上。
import pandas as pd names = ["avg_load", "peak_load", "trough_load"] # 1) 任务相关性:看非对角值,判断是否存在共享结构 corr = pd.DataFrame(np.corrcoef(Y.T), index=names, columns=names) print(corr.round(3)) # 2) 逐任务 RMSE:多任务模型的验收标准 def multi_task_rmse(pred, y_true): return {names[i]: round(float(np.sqrt(np.mean((pred[:, i] - y_true[:, i]) ** 2))), 3) for i in range(y_true.shape[1])}完整案例:预测次日平均负荷、峰值负荷与谷值负荷
场景:一个配电网公司要同时得到次日平均负荷、峰值负荷和谷值负荷,用于排产与调峰。三个目标都由同一组天气、日历特征驱动,是典型的强相关多任务。
先造一份仿真数据:
import numpy as np import lightgbm as lgb from sklearn.multioutput import MultiOutputRegressor from sklearn.model_selection import train_test_split rng = np.random.default_rng(7) n = 2000 temp = rng.normal(20, 8, n) # 气温 weekend = rng.integers(0, 2, n) # 周末标记 X = np.column_stack([temp, weekend]) base = 10 + 0.05 * temp + 0.1 * weekend avg_load = base + rng.normal(0, 0.5, n) peak_load = 1.3 * base + rng.normal(0, 0.6, n) trough_load = 0.8 * base + rng.normal(0, 0.4, n) Y = np.column_stack([avg_load, peak_load, trough_load])下面把两条路径放一起跑,用上一节的multi_task_rmse对比:
Xtr, Xte, Ytr, Yte = train_test_split(X, Y, random_state=42) # 路径一:每个任务独立建模 base_model = MultiOutputRegressor(lgb.LGBMRegressor(n_estimators=300, verbose=-1)) base_model.fit(Xtr, Ytr) pred_base = base_model.predict(Xte) # 路径二:标签堆叠,一个模型同时预测 3 个目标 k = 3 bst = lgb.train({"objective": "regression", "metric": "l2", "verbose": -1}, lgb.Dataset(np.tile(Xtr, (k, 1)), Ytr.T.reshape(-1)), num_boost_round=300) pred_stack = bst.predict(np.tile(Xte, (k, 1))).reshape(k, -1).T print("独立模型:", multi_task_rmse(pred_base, Yte)) print("多任务模型:", multi_task_rmse(pred_stack, Yte))在目标强相关的仿真数据上,两种方案的逐任务 RMSE 通常很接近。这恰恰说明要先看相关性、再做逐任务评估:多任务路线在真实数据上是否占优,不是天然成立的,要用这套评估量出来。
性能调优
训练成本是多任务和单任务最直观的差别,⚡ 抓住三个关键参数即可。
模型间并行,模型内限线程。独立模型方案里各任务互不依赖,进程级并行最划算;同时把每个模型的n_jobs调小,避免线程超额订阅:
from joblib import Parallel, delayed def train_one(i): m = lgb.LGBMRegressor(n_estimators=500, n_jobs=4, random_state=42) m.fit(Xtr, Y[:, i], eval_set=[(Xte, Y[:, i])], eval_metric="l2", callbacks=[lgb.early_stopping(50, first_metric_only=True, verbose=False)]) return m models = Parallel(n_jobs=3)(delayed(train_one)(i) for i in range(Y.shape[1])) # models[0].best_iteration_ 最优轮数;models[0].best_score_ 最优验证得分早停控制轮数。n_estimators给足余量,靠lgb.early_stopping(stopping_rounds)在验证曲线走平时刹车;min_delta能过滤微小抖动,first_metric_only=True表示只认第一个指标。早停对多任务尤其重要——任务多了以后,过训练的浪费是成倍放大的。
大数据量交给 GPU。参数里设device="gpu",同一份数据的训练耗时能明显下降:
上图为官方性能测试中不同硬件与分箱配置的训练耗时对比。多任务场景下数据规模随任务数增长,GPU 的杠杆会被进一步放大。
选型清单:该选哪条路
✅ 按这五条对号入座:
- 任务类型混杂(二分类 + 回归)或相关性低:用
MultiOutputClassifier/MultiOutputRegressor包装,稳定好维护。 - 全是二分类且共用样本集:
objective="multilabel_binary",一个模型通吃。 - 回归目标相关性强、特征集相同:标签堆叠 + 自定义目标函数加权。
- 不同任务特征不一致、或任务数会持续增长:任务标识特征的单一模型。
- 数据量大或轮数多:
device="gpu"+ 早停,把成本压住。
【免费下载链接】LightGBMA fast, distributed, high performance gradient boosting (GBT, GBDT, GBRT, GBM or MART) framework based on decision tree algorithms, used for ranking, classification and many other machine learning tasks.项目地址: https://gitcode.com/GitHub_Trending/li/LightGBM
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考