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.