Multi-Task Learning for Simultaneous Price, Volume, and Volatility Prediction
多任务学习(MTL)通常带有这样的断言:在相关目标之间共享编码器,并且主要任务会变得更好。在交易中,相关目标是显而易见的——回报、交易量和实现的波动性都落在同一个订单流之外——而且这一断言几乎从未经过检验。有趣的问题不在于任务是否相关。问题在于共享梯度是否一致,以及不一致的折叠处会发生什么。
本文将两件事置于中心位置,大多数 MTL 文章都将其视为脚注:
- 损失平衡是实验,而不是细节。 固定权重、Kendall 不确定性加权和 GradNorm 是三种不同的模型。在相同的折叠上运行所有三个,并报告学习到的权重以及每个的主要任务指标。
- 在看到指标之前,负迁移是可以测量的。 共享编码器上的任务梯度之间的余弦相似性告诉您,在训练期间,辅助任务是否将表示拉到主要任务想要去的地方。对余弦进行签名,然后检查该符号是否预测了该折叠的结果。
管道中的其他所有内容——波动性过程、训练循环、泄漏控制、验证协议——已经在本博客的其他地方介绍过,并且是链接的,而不是重新派生的。
## 设置

给定输入特征(OHLCV、技术指标、订单流),三个目标:
- 任务 1(主要):下一期回报
- 任务2(辅助):下一周期日志量
- 任务 3(辅助):下一期已实现波动率
多任务模型同时产生所有三个,,多任务风险是每个任务风险的加权和:
整篇文章都是关于 以及每个任务梯度之间的作用。
为什么联合训练可能会有所帮助,在一段话中。 辅助任务限制共享表征来解释多个市场现象,这同时是一种容量控制和归纳偏差;由于可以直接观察到成交量和波动性,而不能直接观察到“预期收益”,因此辅助头可提供比主头更清晰的梯度信号。在用于多水平预测的时间融合变换器]中,对一个模型发出许多输出的情况进行了详细的论证——附加了可解释性机制,这为多水平分位数提出了相同的共享编码器多头论证。
架构概览
硬参数共享:共享编码器 馈送 任务特定头 ,因此 。这是这里测量的版本,因为它是 上梯度冲突定义明确的版本。
软参数共享为每个任务提供了自己的编码器,并带有耦合惩罚 - 更多参数,更大灵活性,并且没有单个共享参数向量来测量冲突。 十字绣网络位于两者之间,通过每个级别的学习矩阵 混合每个任务的功能。如果硬共享显示出冲突,则两者都值得尝试,并且两者都超出了下面的测量范围。
重要的实验:三种损失平衡方案

朴素损失 对规模敏感。如果回波损耗在 附近,体积损耗在 附近,则体积拥有梯度,并且回波头匮乏。三个回应:
固定权重。 在标准化每个目标后设置 。诚实的基线——如果它赢了,自适应方案就只是仪式。
**不确定性加权(Kendall et al., 2018)。**学习每个任务的同方差噪声量表 :
高不确定性任务会自动降低权重; 项阻止了简单的 解决方案。请注意,这个 是一个训练时间损失加权设备,而不是一个预测区间 - 对于不确定性,您可以实际调整位置,请参阅适形预测。
GradNorm(Chen 等人,2018)。 平衡梯度幅度而不是损失尺度。每一步:计算和均值,计算相对训练率,并更新。然后,无论损失规模如何,所有任务都以相当的速率进行训练。
MTL 特定的代码是头部、列表向前返回和损失聚合。 Linear/BatchNorm/ReLU/Dropout 堆栈、Adam/cosine/clip 样板和 epoch 循环是 DeepLOB 中显示的标准模式,此处省略。
import torch
import torch.nn as nn
class MultiTaskTradingModel(nn.Module):
"""Hard parameter sharing: one encoder, K heads."""
def __init__(self, encoder: nn.Module, repr_dim: int, n_tasks: int = 3):
super().__init__()
self.shared_encoder = encoder # any MLP/CNN/GRU trunk
self.task_heads = nn.ModuleList(
nn.Linear(repr_dim, 1) for _ in range(n_tasks)
)
def forward(self, x):
h = self.shared_encoder(x)
return [head(h).squeeze(-1) for head in self.task_heads]
def shared_repr(self, x):
return self.shared_encoder(x)
class UncertaintyWeightedLoss(nn.Module):
"""Kendall et al. (2018) homoscedastic weighting."""
def __init__(self, n_tasks: int = 3):
super().__init__()
self.log_vars = nn.Parameter(torch.zeros(n_tasks)) # log(sigma^2)
def forward(self, losses: list) -> torch.Tensor:
return sum(
torch.exp(-self.log_vars[i]) * loss + self.log_vars[i]
for i, loss in enumerate(losses)
)
def get_weights(self) -> list:
with torch.no_grad():
return [torch.exp(-lv).item() for lv in self.log_vars]
UncertaintyWeightedLoss 有参数,因此它必须与模型 optim.Adam(list(model.parameters()) + list(uw.parameters()), ...) 一起进入优化器。忘记这一点是“运行不确定性加权”并默默地运行固定权重的最常见方法。
报告内容
对于每个方案,在每个折叠上:学习的最终任务权重、主要任务指标,以及——因为权重方案是一种模型选择——在选择一个方案之前比较了多少个方案。
| 方案 | 主要任务指标与单任务指标 | |||
|---|---|---|---|---|
| 已修复() | 1.00 | 1.00 | 1.00 | — |
| 不确定性加权 | — | — | — | — |
| 研究生规范 | — | — | — | — |
三个方案乘以几倍已经是一个小型的模型搜索了。这里报告的任何改进都必须经过泄气的 Sharpe 和多重测试 中描述的多重测试修正后才有意义。
负迁移:梯度揭示了什么

这是值得保留的部分。负迁移是指辅助任务使主要任务变得更糟,它有一个直接的诊断:共享参数空间中任务梯度之间的角度。
仅在共享编码器上测量 - 头在构造上是特定于任务的,并且总是微不足道地“一致”。
import torch.nn.functional as F
def shared_grad(model, x, y, task_idx, criterion=nn.MSELoss()):
"""Gradient of task `task_idx` w.r.t. the shared encoder, flattened."""
model.zero_grad(set_to_none=True)
loss = criterion(model(x)[task_idx], y)
loss.backward()
return torch.cat([
p.grad.detach().flatten()
for p in model.shared_encoder.parameters()
if p.grad is not None
])
def task_conflict(model, x, y_by_task, task_names):
"""Pairwise cosine similarity between per-task shared-encoder gradients."""
grads = {
name: shared_grad(model, x, y_by_task[name], i)
for i, name in enumerate(task_names)
}
return {
(a, b): F.cosine_similarity(
grads[a].unsqueeze(0), grads[b].unsqueeze(0)
).item()
for i, a in enumerate(task_names)
for b in task_names[i + 1:]
}
在训练期间以固定的节奏对保留的批次调用此方法,而不是在结束时调用此方法。当编码器专门化时,一对可以开始对齐和发散;一个单一的训练结束数字就隐藏了这一点。
寻找并以任一方式发布的发现:
| 配对 | cos sim,早期训练 | cos sim,后期训练 | MTL 对首要任务有帮助吗? |
|---|---|---|---|
| 返回 ↔ 音量 | — | — | — |
| 回报率 ↔ 波动率 | — | — | — |
| 交易量 ↔ 波动性 | — | — | — |
如果交易量和波动率梯度彼此一致,同时又与回报梯度相冲突,则正确的结论是,这两个辅助任务形成了回报任务不属于的连贯块,而解决办法是任务分组,而不是增加容量。当冲突确实存在时,标准补救措施是PCGrad(Yu et al., 2020),它将每个冲突的梯度投影到另一个的法线平面上; CAGrad (Liu et al., 2021),它搜索不损害任何任务的下降方向;或者完全放弃辅助任务。
请注意故意缺失的内容:按目标值着色的共享表示的 t-SNE 图。它是装饰性的——上面的余弦数字说明了嵌入所表示的一切,并且它们将其表示为数字。
验证协议

在草率的协议下,上述测量毫无价值,而 MTL 使通常的陷阱变得更糟,因为有 3 个目标要泄漏,而不是 1 个。
真实数据,而不是模拟器。 目标必须来自实际的 OHLCV/贸易数据。硬编码的 GARCH 玩具产生的波动性与通过构造的回报相关,这正是被测试的东西——实验将测量它自己的生成器。如果您想要一个拟合的波动率过程,crypto 的 GARCH 波动率预测通过真实 BTC/ETH 上的最大似然拟合 GARCH(1,1) 并验证标准化残差,并且非对称 GARCH 和杠杆效应 涵盖了为什么高斯对称响应模拟器首先会错误地预测加密波动率。仅当合成数据提供受控的基本事实(您试图恢复的已知的、作者设置的任务相关性)时,它才是有道理的,这是与此处的实验不同的实验。
缩放器仅适合训练。 将特征缩放器和所有三个目标缩放器安装在每个训练折叠内并应用于验证;在将泄漏测试集时刻分解为训练之前的全局 fit_transform。这个确切的失败被分类在前瞻偏差分类法中。
清除、禁运的前向折叠。 一个 80/20 时间顺序分割无法区分 MTL 改进和折叠效果 - 这就是 前向优化 的整个论点,它显示了三个分割产生三个结论。重用使用机器学习进行扩展建模]中的扩展窗口 purged_walk_forward 生成器:它会在每个边界两侧删除 horizon 行的间隙,这在这里很重要,因为即使返回目标没有泄漏,重叠的已实现波动率窗口也会跨越边界泄漏。
经典基线。 如果每个目标梯度提升或岭模型击败所有四个网络,则击败三个单任务网络的 MTL 网络并不能证明什么。在相同的折叠和相同的特征上使用 LightGBM 或脊为每个目标拟合一个模型,并将其报告在同一个表中。
| 型号 | 主要任务指标 | 笔记 |
|---|---|---|
| 山脊,每个目标 | — | 经典基线 |
| LightGBM,每个目标 | — | 经典基线 |
| 每个目标的单任务 MLP | — | 三个独立的网 |
| MTL,最佳损失方案 | — | 一网三头 |
MTL 在什么情况下值得采用

MTL 应该获胜的条件,以假设的形式陈述,以检查上述折叠,而不是作为清单:
- 辅助标签比主标签更干净。直接观察体积; “预期回报”不是。如果返回头主要是拟合噪声,则来自辅助头的梯度信号是物镜的唯一适定部分。
- 训练数据相对于编码器容量是有限的,因此辅助约束确实进行正则化工作,而不仅仅是竞争参数。
- 推理延迟很重要,一次前向传递胜过三次。
反对的情况同样可测试:如果测量的 cos_sim(return, ·) 值持续为负,则共享编码器将脱离主要任务,并且辅助头是一种负担,而不是正则化器。
## 结论

回报、交易量和波动性来自相同的微观结构,因此共享表示是合理的先验,但先验并不是结果。这个设置实际上可以建立的两件事是数据更喜欢哪种损失平衡方案(报告学习的权重,而不仅仅是命名的获胜者)以及共享编码器上的任务梯度是否一致,通过训练进行测量,而不是根据目标相关的事实进行假设。
如果清除的前向折叠显示 MTL 网络未能击败每个目标的梯度提升模型,这就是发现,并且它会这样发布 - 模板是诚实的负 。负迁移的负结果仍然是负迁移的结果。
Authors
Trading-systems engineer
Trading-systems engineer building bots since 2017: cross-exchange arbitrage (connected up to 30 venues), cointegration-based pairs arbitrage across spot and futures, scalping, news and sentiment-driven strategies, trend algorithms, and portfolio management and balancing algorithms. Also builds sub-millisecond order execution, big-data warehouses, backtesting engines, AI agents, and trading interfaces (incl. open-source profitmaker.cc). Stack: JS/TS, Python, Rust/Zig/Go, DevOps, backend, frontend, architecture.