Replace the serial curiosity_fwd_bwd_per_block kernel (CUR_TOTAL_PARAMS=11306 loop iterations × block-level shared-memory reduction per-iteration = 1085ms) with a cuBLAS GEMM pipeline matching the existing curiosity inference path: Forward: curiosity_prepare_input → GEMM1(W1) → bias_leaky_relu → GEMM2(W2) → mse_fwd_grad Backward: gemm_dw(dW2) → bias_grad_reduce(db2) → gemm_dx(d_hidden) → leaky_relu_bwd → gemm_dw(dW1) → bias_grad_reduce(db1) New CUDA kernels added to curiosity_training_kernel.cu: - curiosity_mse_fwd_grad: +b2 in-place, d_pred = 2/CUR_OUTPUT*(pred-target) - curiosity_leaky_relu_bwd: gates d_hidden by sign of post-activation hidden - curiosity_bias_grad_reduce: sum dy[N, D] over batch → grad_b[D] GpuCuriosityTrainer rewritten with CuriosityGemm (dedicated cuBLAS+cublasLt handle) + intermediate buffers (input_buf, hidden_buf, pred_buf, d_hidden_buf). Reuses forward kernels from curiosity_inference_kernel.cu. Keeps curiosity_adam_step. Drops partial_grads buffer (max_blocks*11306 floats saved). Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
ml
10-model ML ensemble for the Foxhunt HFT system, built on Candle v0.9.1.
Models
- DQN (Rainbow) — deep Q-network with prioritized replay, dueling heads, noisy nets
- PPO — proximal policy optimization with GAE, LSTM policies, clip-higher
- TFT — temporal fusion transformer for multi-horizon forecasting
- Mamba2 — state space model for sequence prediction
- Liquid Networks — biologically inspired networks for non-stationary data
- TLOB — transformer-based limit order book analysis
- KAN — Kolmogorov-Arnold networks
- xLSTM — extended LSTM architecture
- TGGN — temporal graph neural network
- Diffusion — diffusion-based generative model
Key Modules
ensemble— model ensemble coordination and confidence aggregationhyperopt— PSO-based hyperparameter optimization with per-model adapterstrainers— unified training loops (DQN, PPO, supervised)inference—InferenceAdaptertrait for predictioncheckpoint— model checkpointing and restorationevaluation— walk-forward evaluation pipeline
Usage
use ml::dqn::DQN;
use ml::ppo::PpoTrainer;