Knowledge Distillation: Compressing Trading Models for Low-Latency Deployment
Ketegangan accuracy-vs-latency dalam dagangan berasaskan ML sudah mempunyai jawapan yang diterbitkan di blog ini. Pemodelan spread dengan machine learning mencadangkan pembahagian dua peringkat: model gradient-boosting yang pantas melakukan quoting masa nyata yang kritikal terhadap latensi, manakala model mendalam berjalan secara asynchronous dan membekalkan signal sekunder atau melaraskan parameternya. Dua model, dua jam, satu sistem.
Knowledge distillation ialah jawapan berbeza kepada ketegangan yang sama. Daripada menjalankan model perlahan bersama model pantas, gunakan model perlahan itu sekali secara offline untuk melatih model pantas — student mempelajari keseluruhan taburan kebarangkalian teacher terhadap hasil, bukan sekadar label keras, kemudian teacher dikeluarkan sepenuhnya daripada hot path. Satu model pada masa inference, tiada coupling asynchronous, tiada staleness window.
Jawapan yang menang ialah soalan empirikal, dan artikel ini belum menjawabnya. Yang berikut ialah machinery serta pernyataan jelas tentang ukuran yang akan menentukannya. Tiada apa-apa di sini merupakan benchmark result; di tempat yang biasanya memerlukan nombor, terdapat penanda yang menyatakan apa yang perlu dijalankan.
Satu pembetulan framing di awal, daripada DeepLOB dan deep learning pada order book: ketepatan klasifikasi yang tinggi tidak terus bermakna profit — pergerakan yang diramal mesti melepasi bid-ask spread. Oleh itu, "mengekalkan ketepatan arah teacher" ialah sasaran yang salah untuk mengoptimumkan setup distillation.
Rangka Kerja Teacher-Student

Formulasi asal oleh Hinton, Vinyals dan Dean (2015) adalah mudah. Anda mempunyai model teacher (besar, perlahan, tepat) dan model student (kecil, pantas, untuk dilatih). Student belajar daripada dua signal serentak:
- Sasaran keras: label ground-truth (contohnya, harga naik atau turun)
- Sasaran lembut: taburan kebarangkalian output teacher untuk semua kelas
Fungsi loss student menggabungkan kedua-duanya:
dengan dan ialah logits teacher dan student, ialah fungsi softmax, ialah parameter temperature, dan mengawal imbangan antara dua komponen loss.
Mengapa Soft Targets Penting untuk Dagangan
Formulasi tiga kelas up/stationary/down untuk mid-price, thresholding , dan sebab imbalance yang terhasil menyebabkan anda melaporkan weighted F1 dan bukannya accuracy semuanya sudah disediakan dalam DeepLOB — gunakan label scheme yang sama di sini. Perkara khusus kepada distillation ialah apa yang teacher keluarkan sebelum argmax: "up" keras membawa satu bit, manakala 0.72/0.21/0.07 turut mengatakan pergerakan itu mungkin terhenti dan hampir pasti tidak akan berbalik. Struktur merentas kelas itu ialah signal latihan tambahan, dan sebab itulah student dengan soft target boleh menggeneralisasi lebih baik daripada student yang sama tetapi dilatih dengan label sahaja.
Amaran tentang apa yang bukan confidence itu. Output softmax bukan ketidakpastian yang calibrated, dan menganggap 0.55 berbanding 0.85 sebagai input position-sizing ialah jalan pintas yang conformal prediction untuk dagangan sengaja menolak — ia mendapatkan sizing daripada lebar interval, edge ratio dan no-trade filter apabila interval merentasi sifar; softmax mentah tidak memberikan mana-mana daripadanya. Untuk membuktikan tuntutan sizing di sini, calibration student perlu diukur berbanding calibration teacher (reliability diagram, ECE) dan ditunjukkan bahawa distillation mengekalkannya. Result itu belum ada dalam artikel ini.
Temperature dan Soft Targets

Parameter temperature mengawal "kelembutan" taburan kebarangkalian. Dengan logits , softmax dengan temperature ialah:
Apabila (softmax standard), taburan itu peaky — kelas dominan menerima sebahagian besar probability mass. Apabila $T meningkat, taburan menjadi lebih rata dan magnitud relatif logits kelihatan dengan lebih jelas.
| Temperature | Kesan | Kes penggunaan |
|---|---|---|
| Softmax standard, peaky | Inference biasa | |
| Pelembutan sederhana | Distillation umum | |
| Pelembutan berat | Apabila teacher sangat yakin | |
| Hampir seragam | Jarang berguna, menghilangkan signal |
Terdapat hujah yang munasabah bahawa model dagangan memerlukan temperature sederhana: ramalan kewangan jauh kurang yakin berbanding klasifikasi imej, jadi teacher mungkin mengeluarkan 0.55/0.30/0.15 dan bukannya 0.99/0.005/0.005, lalu kurang peakiness untuk dilembutkan sebelum signal hilang. Itu hujah, bukan dapatan — julatnya mesti datang daripada sweep pada data sebenar, dinilai dengan weighted F1, dan mungkin berbeza mengikut regime.
Faktor dalam istilah KL divergence mengimbangi magnitud gradient yang berkurang pada temperature lebih tinggi. Tanpanya, loss distillation akan menjadi sangat kecil apabila meningkat.
Memilih Temperature melalui Grid Search
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
from sklearn.metrics import f1_score
def distillation_loss(
student_logits: torch.Tensor,
teacher_logits: torch.Tensor,
labels: torch.Tensor,
temperature: float,
alpha: float,
) -> torch.Tensor:
"""Combined hard-target + soft-target distillation loss."""
hard_loss = F.cross_entropy(student_logits, labels)
soft_teacher = F.log_softmax(teacher_logits / temperature, dim=-1)
soft_student = F.log_softmax(student_logits / temperature, dim=-1)
soft_loss = F.kl_div(
soft_student,
soft_teacher,
log_target=True,
reduction="batchmean",
)
return alpha * hard_loss + (1.0 - alpha) * (temperature ** 2) * soft_loss
def search_temperature(
teacher: nn.Module,
student_factory, # callable returning a fresh student
train_loader: DataLoader,
val_loader: DataLoader,
temperatures: list[float] = [1, 2, 3, 5, 8, 12],
alpha: float = 0.3,
epochs: int = 30,
lr: float = 1e-3,
device: str = "cuda",
):
"""Grid search over temperature, scored by weighted F1 (not accuracy:
the up/flat/down label scheme is heavily imbalanced toward flat)."""
best_f1, best_T, best_student = 0.0, 1.0, None
for T in temperatures:
student = student_factory().to(device)
optimizer = torch.optim.AdamW(student.parameters(), lr=lr)
for epoch in range(epochs):
student.train()
for X, y in train_loader:
X, y = X.to(device), y.to(device)
with torch.no_grad():
teacher_logits = teacher(X)
student_logits = student(X)
loss = distillation_loss(
student_logits, teacher_logits, y, T, alpha
)
optimizer.zero_grad()
loss.backward()
optimizer.step()
student.eval()
preds, targets = [], []
with torch.no_grad():
for X, y in val_loader:
preds.append(student(X.to(device)).argmax(dim=-1).cpu())
targets.append(y)
f1 = f1_score(
torch.cat(targets), torch.cat(preds), average="weighted"
)
print(f"T={T:>4.1f} val_weighted_f1={f1:.4f}")
if f1 > best_f1:
best_f1, best_T, best_student = f1, T, student
print(f"\nBest temperature: T={best_T}, val_weighted_f1={best_f1:.4f}")
return best_T, best_student
Mendistil Ensembles menjadi Satu Model

Quant ensemble menggabungkan inductive bias: gradient-boosted tree pada ciri order-book, 1D-CNN pada tick terkini, transformer pada window berbilang timeframe, dan model linear pada faktor makro. Averaging lebih stabil daripada mana-mana member secara bersendirian, dan menjalankan keempat-empatnya menggandakan latency serta kos — keadaan yang dikendalikan oleh two-stage split daripada pemodelan spread dengan machine learning dengan menurunkan member perlahan ke side channel asynchronous. Sebaliknya, distillation meruntuhkan keempat-empatnya menjadi satu student dalam hot path.
Output teacher ensemble ialah purata output softmax member-membernya:
dengan ialah bilangan member ensemble. Student dilatih terhadap taburan purata ini.
class EnsembleTeacher(nn.Module):
"""Wraps K models, returns averaged logits for distillation."""
def __init__(self, models: list[nn.Module]):
super().__init__()
self.models = nn.ModuleList(models)
@torch.no_grad()
def forward(self, x: torch.Tensor) -> torch.Tensor:
logits = torch.stack([m(x) for m in self.models], dim=0)
return logits.mean(dim=0) # average logits, not softmax
class TradingStudent(nn.Module):
"""Lightweight MLP for sub-millisecond inference."""
def __init__(self, input_dim: int, hidden: int = 64, n_classes: int = 3):
super().__init__()
self.net = nn.Sequential(
nn.Linear(input_dim, hidden),
nn.ReLU(),
nn.BatchNorm1d(hidden),
nn.Linear(hidden, hidden),
nn.ReLU(),
nn.BatchNorm1d(hidden),
nn.Linear(hidden, n_classes),
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.net(x)
Asimetri kiraan parameter ialah keseluruhan tujuannya: MLP dua lapisan dengan 64 hidden units mempunyai kira-kira 8,000 parameter untuk tugasan 60-ciri, 3-kelas, berbanding kiraan gabungan ensemble yang mencecah jutaan.
Apa yang Dikekalkan dan Hilang oleh Student
Ini ialah soalan empirikal yang paling penting dan artikel ini tidak menjawabnya. Intuisinya ialah student mengekori ensemble dalam distribution yang sama dan merosot dalam regime tertekan, apabila kepelbagaian ensemble melakukan tugasnya — tetapi angka retention hanya bermakna jika diukur pada data order-book sebenar, dipecahkan mengikut regime dan dilaporkan sebagai weighted F1. Student yang bertahan pada hari tenang tetapi runtuh ketika liquidation cascade ialah produk yang berbeza daripada student yang merosot secara beransur-ansur, dan nombor aggregate tidak dapat membezakan kedua-duanya.
Tiga mitigation patut diuji berdasarkan ukuran itu, bukan didakwa terlebih dahulu:
- Sertakan tempoh tertekan dalam set distillation supaya student melihat regime ketika jurang dijangka terbuka.
- Feature-based distillation — padankan representation perantaraan, bukan output akhir sahaja.
- Auxiliary regime head pada student, yang memaksa ciri sedar-regime masuk ke shared trunk.
Self-Distillation: Apabila Student Menjadi Teacher

Self-distillation ialah teknik apabila model mendistilkan pengetahuan daripada dirinya sendiri.
Born-Again Networks (BANs)
Latih student dengan architecture yang sama seperti teacher. Student "born-again" sering mengatasi model asal, dan prosesnya diulang:
Setiap generation dilatih menggunakan soft targets daripada generation sebelumnya, dengan gain biasanya tepu selepas beberapa generation. Untuk model dagangan, kos architecture ialah sifar — tiada feature baharu, tiada data baharu, hanya prosedur latihan yang berbeza — maka murah untuk diuji dan tiada alasan untuk melaporkannya tanpa ujian.
Depth-Wise Self-Distillation
Pasangkan auxiliary classifier pada lapisan perantaraan. Exit paling dalam menjadi teacher untuk exit yang lebih cetek. Semasa inference, pilih exit: cetek untuk latency lebih rendah, dalam untuk accuracy maksimum.
Inilah idea yang paling sesuai dengan sistem dagangan, kerana kedalaman exit menjadi tombol latency masa runtime: satu network yang dilatih meliputi pelbagai budget dan bukannya memaksa satu architecture ketika training. Apabila book bergerak pantas, gunakan exit cetek dan terima posterior yang lebih buruk; apabila keadaan tenang, bayar kos full depth. Curve accuracy-per-exit dan latency-per-exit boleh diukur, dan crossovernya menentukan sama ada tombol itu berbaloi.
class SelfDistillingNet(nn.Module):
"""Network with early-exit classifiers for variable-latency inference."""
def __init__(self, input_dim: int, n_classes: int = 3):
super().__init__()
self.block1 = nn.Sequential(
nn.Linear(input_dim, 128), nn.ReLU(), nn.BatchNorm1d(128)
)
self.block2 = nn.Sequential(
nn.Linear(128, 64), nn.ReLU(), nn.BatchNorm1d(64)
)
self.block3 = nn.Sequential(
nn.Linear(64, 32), nn.ReLU(), nn.BatchNorm1d(32)
)
self.exit1 = nn.Linear(128, n_classes)
self.exit2 = nn.Linear(64, n_classes)
self.exit3 = nn.Linear(32, n_classes) # final exit
def forward(
self, x: torch.Tensor, exit_layer: int = 3
) -> torch.Tensor:
h1 = self.block1(x)
if exit_layer == 1:
return self.exit1(h1)
h2 = self.block2(h1)
if exit_layer == 2:
return self.exit2(h2)
h3 = self.block3(h2)
return self.exit3(h3)
def forward_all_exits(self, x: torch.Tensor):
"""Return logits from all exits (for self-distillation training)."""
h1 = self.block1(x)
h2 = self.block2(h1)
h3 = self.block3(h2)
return self.exit1(h1), self.exit2(h2), self.exit3(h3)
def self_distillation_step(
model: SelfDistillingNet,
x: torch.Tensor,
y: torch.Tensor,
temperature: float = 4.0,
alpha: float = 0.5,
) -> torch.Tensor:
"""One training step with self-distillation from deepest exit."""
logits_1, logits_2, logits_3 = model.forward_all_exits(x)
loss_hard = F.cross_entropy(logits_3, y)
loss_distill_1 = distillation_loss(
logits_1, logits_3.detach(), y, temperature, alpha
)
loss_distill_2 = distillation_loss(
logits_2, logits_3.detach(), y, temperature, alpha
)
return loss_hard + 0.5 * loss_distill_1 + 0.5 * loss_distill_2
Dari Mana Datangnya Inference Budget

Distillation hanya penting jika inference perlu berada dalam budget yang keras, dan keseluruhan tick-to-trade ladder — NIC-to-userspace, kernel bypass, jumlah sub-100 µs, serta tier sub-10 µs yang memaksa penggunaan FPGA dan shared memory — sudah diterangkan dalam data dan komunikasi dalam algorithmic trading. Baris yang masih terbuka dalam ladder itu ialah model inference, dan itulah baris yang cuba diisi oleh distillation.
Jangan mengisi baris lain dengan jadual latency mengikut kelas model. Pemodelan spread dengan machine learning sudah menerbitkan perbandingan GBM-vs-deep-learning serta caveat yang lebih penting daripada nombor: latency bergantung pada implementation, dan model LightGBM yang sama mengambil puluhan microseconds setiap baris daripada Python tetapi hanya beberapa microseconds daripada predictor yang compiled. Sebarang tuntutan latency di sini mesti menamakan framework, core dan batch size, jika tidak ia hanya noise.
Khusus pada GPU: overhead tetap bagi setiap launch perlu diamortisasi dahulu sebelum device memberi manfaat, dan inference single-row berada jauh di sebelah kiri roofline ridge, tempat ia tidak pernah sampai. Bila GPU berbaloi mengukur curve amortisasi itu dengan betul melalui batch sweep, termasuk bagaimana kad PCIe diskret menolak ridge lebih jauh ke kanan — baca itu dan jangan percaya constant yang dipetik daripada ingatan.
Quantization selepas Distillation
Student yang didistil boleh dimampatkan lagi: weights INT8 (kira-kira 2x pada CPU dengan AVX-512 VNNI), weights binary/ternary yang menukar darab kepada tambah, dan pruning untuk melangkau pengiraan hampir sifar.
Dakwaan yang menarik ialah distillation-then-quantization mengekalkan lebih banyak accuracy berbanding quantization sahaja, kerana student sudah mempelajari representation yang kompak. Jangan deploy berdasarkan dakwaan itu. GPU precision trap ialah pendirian blog ini terhadap reduced numeric precision: ia secara senyap memulangkan garbage yang kelihatan munasabah, dan perkara yang menjadikan fast path boleh dihantar ialah equivalence gate yang diukur — fills yang berubah, PnL delta dalam bps — bukan assertion. Student INT8 ialah model yang berbeza sehingga gate itu diukur terhadap student FP32.
import torch.quantization as quant
def quantize_student(student: nn.Module, calibration_loader: DataLoader):
"""Post-training static quantization for CPU deployment."""
student.cpu()
student.eval()
student.qconfig = quant.get_default_qconfig("x86")
student_prepared = quant.prepare(student)
with torch.no_grad():
for X, _ in calibration_loader:
student_prepared(X)
student_quantized = quant.convert(student_prepared)
return student_quantized
FPGA Deployment: Pipeline Distill-to-Bitstream

FPGA ialah tier sub-10 µs dalam latency ladder, dan ulasan Tbricks/Broadridge membincangkannya dalam production bersama kernel-bypass NIC — latency deterministic, tiada OS jitter, ditempatkan bersama network stack. Perkara yang tidak dibincangkan di mana-mana dalam blog ini ialah cara model yang didistilkan dimuatkan ke dalamnya.
Nota production DeepLOB menyenaraikan ONNX/TensorRT, quantization INT8 dan deployment FPGA sebagai tiga pilihan lalu berhenti di situ. Pilihan ketiga berkembang seperti berikut:
1. Train ensemble teacher (offline, GPU cluster, hours/days)
|
2. Distill to small MLP student (offline, single GPU, minutes)
|
3. Quantize student to INT8 / fixed-point (offline, CPU)
|
4. Convert to HLS (High-Level Synthesis) or RTL
|
5. Synthesize FPGA bitstream (offline, hours)
|
6. Deploy to FPGA card in production server
|
7. Inference: market data -> FPGA -> trading signal
Kekangan utama ialah model mesti muat dalam logic elements device — LUT, DSP slice dan block RAM. Sebagai budget order-of-magnitude dan bukannya ukuran: MLP 2 lapisan dengan 64 hidden units dan weights INT8 memerlukan kira-kira 8,000 multiply-accumulates bagi setiap inference dan sekitar 16 KB weights, sebahagian kecil daripada part kelas pertengahan. Di sinilah distillation berbaloi — teacher ensemble tidak muat pada mana-mana budget; student jauh daripada had itu.
Tools yang mengautomatikkan PyTorch/ONNX menjadi hardware yang boleh disintesis termasuk AMD/Xilinx Vitis AI, hls4ml (daripada CERN) dan FINN (daripada Xilinx Research).
Contoh: Penukaran hls4ml
import hls4ml
import onnx
dummy_input = torch.randn(1, 60) # 60 input features
torch.onnx.export(student, dummy_input, "student.onnx", opset_version=13)
hls_config = hls4ml.utils.config_from_onnx_model(
onnx.load("student.onnx"),
granularity="name",
default_precision="ap_fixed<16,8>",
default_reuse_factor=1, # full parallelism
)
hls_model = hls4ml.converters.convert_from_onnx_model(
"student.onnx",
hls_config=hls_config,
output_dir="hls_student",
backend="VivadoAccelerator",
board="alveo-u250",
)
hls_model.compile()
hls_model.build(csim=True, synth=True)
hls_model.report()
hls_model.report() ialah satu-satunya sumber yang boleh dipercayai untuk nombor resource dan latency bagi model, board, precision dan reuse factor tertentu — angka berubah banyak hanya dengan default_reuse_factor. Memetik jadual sintesis "tipikal" tanpa menjalankannya hanyalah meneka.
Pertimbangan Praktikal

Precomputing Teacher Logits
Distillation memerlukan prediction teacher untuk keseluruhan training set — kos offline sekali sahaja yang wajar dibayar dengan sengaja: jalankan ensemble sekali, simpan logits, dan latih student terhadap cache. Temperature sweep dan architecture search selepas itu tidak memerlukan forward pass teacher tambahan, dan itulah yang menjadikan sweep di atas praktikal.
Monitor Khusus Distillation
Kebersihan feature pipeline, rolling normalization kerana parameter z-score drift, pemantauan input distribution-shift dan retraining yang dicetuskan regime semuanya dibincangkan dalam production section DeepLOB dan terpakai di sini tanpa perubahan.
Monitor khusus untuk distillation ialah teacher-student KL divergence pada data live. Teacher masih wujud secara offline; jalankannya pada sampel input live dan bandingkan taburan. KL yang meningkat bermakna approximation student merosot dalam regime yang tidak digunakan untuk distillation — dan ia tercetus sebelum accuracy, kerana tidak menunggu label. Threshold retraining perlu calibrated terhadap KL yang diperhatikan dalam tempoh yang diketahui baik dan merosot; jika dipilih a priori, ia arbitrari.
Bila Tidak Patut Distill
- Teacher sudah kecil (model linear, GBM cetek): distillation menambah satu peringkat pipeline tanpa compression.
- Latency bukan kekangan (rebalancing harian, signal hujung hari): deploy teacher.
- Interpretability mengatasi speed: network yang didistil lebih sukar diterangkan daripada ensemble tree yang digantikannya.
- Two-stage split sudah berfungsi: jika model slow asynchronous dalam architecture pemodelan spread memberikan hasil, distillation perlu mengatasinya dalam perbandingan yang diukur sebelum menggantikan sistem yang berfungsi.
Ringkasan

Distillation ialah alternatif yang koheren kepada two-stage fast/slow split: latih teacher terbaik yang mampu anda biayai secara offline, pindahkan struktur soft-targetnya ke student yang cukup kecil untuk hot path, quantize, kemudian deploy pada CPU atau FPGA. Variant depth-wise melangkah lebih jauh dan menjadikan latency pilihan runtime dan bukannya pilihan masa training.
Apa yang sengaja tidak didakwa oleh artikel ini ialah bahawa mana-mana daripadanya mengatasi apa yang telah diterbitkan oleh blog ini. Keputusan itu memerlukan tiga ukuran pada data order-book sebenar: curve retention weighted F1 student-vs-ensemble yang dipecahkan mengikut regime, temperature sweep dan parity gate INT8 mengikut gaya GPU precision trap. Sehingga semua itu wujud, ini ialah penerangan tentang teknik, bukan cadangan untuk deploy.
Rujukan
-
Hinton, G., Vinyals, O., & Dean, J. (2015). Distilling the Knowledge in a Neural Network. arXiv:1503.02531
-
Furlanello, T., Lipton, Z. C., Tschannen, M., Itti, L., & Anandkumar, A. (2018). Born-Again Neural Networks. ICML. arXiv:1805.04770
-
Zhang, L., Song, J., Gao, A., Chen, J., Bao, C., & Ma, K. (2019). Be Your Own Teacher: Improve the Performance of Convolutional Neural Networks via Self Distillation. ICCV. arXiv:1905.08094
-
Romero, A., Ballas, N., Kahou, S. E., Chassang, A., Gatta, C., & Bengio, Y. (2015). FitNets: Hints for Thin Deep Nets. ICLR. arXiv:1412.6550
-
Gou, J., Yu, B., Maybank, S. J., & Tao, D. (2021). Knowledge Distillation: A Survey. International Journal of Computer Vision, 129, 1789-1819. arXiv:2006.05525
-
Duarte, J., et al. (2018). Fast Inference of Deep Neural Networks in FPGAs for Particle Physics (hls4ml). Journal of Instrumentation, 13, P07027. arXiv:1804.06913
-
Umuroglu, Y., et al. (2017). FINN: A Framework for Fast, Scalable Binarized Neural Network Inference. FPGA '17. arXiv:1612.07119
-
Zhang, Z., Zohren, S., & Roberts, S. (2019). DeepLOB: Deep Convolutional Neural Networks for Limit Order Books. IEEE Transactions on Signal Processing, 67(11), 3001-3012. arXiv:1808.03668
Pengarang
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.