spec(ml-alpha): per-horizon attention pool design (C20)
Captures the brainstormed "alternative attention pool variants"
follow-on from the original real-LOB integration brainstorm (Axis 1,
deferred from the LOB workstream as a separate model-side spec).
Design:
Replace shared learned query Q[HIDDEN_DIM] with per-horizon queries
Q_h[N_HORIZONS, HIDDEN_DIM]. Per-horizon softmax + context vectors
feed multi-horizon heads directly (PATH A) — each horizon attends
to a different part of the K=6000 LN_b output sequence. CfC k=0
state is initialised by the MEAN of per-horizon contexts so the
K-loop recurrence + state amplification (per
pearl_state_amplifies_short_horizon_into_long_horizon) survives.
Heads consume per-horizon context concat CfC h_K (residual) with a
default 75/25 weight split.
Falsifiable claim (§0): A/B-tested win means h6000 mean_auc lifts by
≥ +0.01 absolute OR per-horizon distribution shifts toward short
horizons (h1000, h300) with no net h6000 loss. The 3-fold variance
band on the current architecture (mean_auc 0.7749 ± 0.024) means a
+0.01 lift is within noise — a meaningful effect needs ≥ +0.024 or
qualitative distributional shift.
Two new kernels (per_horizon_attention_pool_fwd + _bwd) + signature
extension on multi_horizon_heads_{fwd,bwd}. Variant-toggle config flag
(SharedQuery vs PerHorizonQuery) keeps the existing path fully
functional; new variant is opt-in. CheckpointV1 → V2 with explicit
discriminant + optional q_h field; V1 files load as SharedQuery, new
V2 training writes the discriminant.
Three validation rings:
1. Per-(b,h) numgrad parity at K=16
2. One-epoch smoke (no NaN, loss decreases)
3. 30-epoch × 3-fold A/B (#204) — decision gate per §0 falsifiable claim
Implementation explicitly deferred. The decision to invest depends on
(a) GPU time budget (~3-6 hrs on L40S × 5 GPUs for the A/B), (b)
whether per-horizon cost-frontier sweeps (#202 follow-ups) surface
viable horizons beyond h6000 that would benefit from per-horizon
specialisation, and (c) the 3-fold variance noise floor making the
expected effect size visible.
Next step when ready: invoke superpowers:writing-plans against this
spec for the ~6-8 commit implementation plan.
Closes the "good to have" question from the recent brainstorm with a
concrete decision framework rather than ad-hoc implementation.
Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,204 @@
|
||||
# Per-Horizon Attention Pool — Design Spec
|
||||
|
||||
**Status:** DESIGN
|
||||
**Date:** 2026-05-18
|
||||
**Branch (future):** `ml-alpha-phase-c` (separate from `ml-alpha-phase-a` real-LOB work)
|
||||
**Owner:** ml-alpha team
|
||||
**Tracked tasks:** #203 (kernels + wiring) + #204 (A/B sweep)
|
||||
|
||||
## 0. Context
|
||||
|
||||
`pearl_phase1d2_decisive_pass` (2026-05-15) validated the Phase 1+2+3 architecture at K=6000 with a **single shared attention query** `Q[HIDDEN_DIM]` pooling LN_b's `[B, K, HIDDEN_DIM]` output into one context vector that initialises the CfC k=0 state. Phase B (`ml-alpha-phase-a`) refactored the kernel to block-per-batch and 3-fold CV at `mean_auc = 0.7749 ± 0.024 / h6000 0.7591 ± 0.018`. The original brainstorm called out per-horizon attention pool variants as a "separate model-side spec".
|
||||
|
||||
**Hypothesis:** different horizons attend to different parts of the K-sequence. A buy-now (h30) signal lives in the most-recent 5-10 snapshots; a 2.4hr (h6000) signal lives in regime-defining state from the entire K window. One shared query has to compromise; per-horizon queries can specialise.
|
||||
|
||||
**Falsifiable claim:** per-horizon attention pool lifts h6000 val AUC by ≥ +0.01 absolute (within the noise floor of single-fold AUC variance per `pearl_single_window_oos_is_not_oos.md`) on a held-out 3-fold CV, OR — at minimum — shifts the per-horizon AUC distribution so short horizons (h30, h100) gain at least as much as h6000 loses. The decision criterion is total cross-fold mean AUC averaged over horizons weighted by Phase-1d.4 deployment economics (currently h6000 = 1.0, others = 0.0; will broaden as cost-frontier sweeps surface other viable horizons).
|
||||
|
||||
Out of scope:
|
||||
- Multi-head attention (more than per-horizon — separate spec).
|
||||
- Cross-attention to external context (no-op for single-instrument ES).
|
||||
- Position embeddings for K (orthogonal — could land after per-horizon).
|
||||
|
||||
## 1. Architecture
|
||||
|
||||
LN_b output: `[B, K, HIDDEN_DIM]` (unchanged).
|
||||
|
||||
**Current** (single query):
|
||||
```
|
||||
Q: [HIDDEN_DIM] shared learned param
|
||||
scores[b, k] = Σ_h Q[h] * LNb[b, k, h] [B, K]
|
||||
attn[b, k] = softmax_k(scores[b]) [B, K]
|
||||
context[b, h] = Σ_k attn[b, k] * LNb[b, k, h] [B, HIDDEN_DIM]
|
||||
```
|
||||
context → CfC k=0 state → K-loop recurrence → multi-horizon heads.
|
||||
|
||||
**Proposed** (per-horizon queries):
|
||||
```
|
||||
Q_h: [N_HORIZONS, HIDDEN_DIM] five learned queries
|
||||
scores_h[b, h, k] = Σ_d Q_h[h, d] * LNb[b, k, d] [B, N_HORIZONS, K]
|
||||
attn_h[b, h, k] = softmax_k(scores_h[b, h]) [B, N_HORIZONS, K]
|
||||
context_h[b, h, d] = Σ_k attn_h[b, h, k] * LNb[b, k, d] [B, N_HORIZONS, HIDDEN_DIM]
|
||||
```
|
||||
|
||||
Two consumption paths to pick from in the spec-review:
|
||||
|
||||
### Path A (recommended) — per-horizon contexts feed heads directly, CfC unchanged
|
||||
- The CfC k=0 state is still initialised by the **mean of per-horizon contexts** (preserves the existing K-loop recurrent semantics without disruption).
|
||||
- The multi-horizon heads kernel changes to read each horizon's own `context_h[b, h, :]` as the *direct* input (bypassing the CfC's final h_K for the head's primary signal).
|
||||
- Heads still consume CfC h_K as a residual / auxiliary input (concatenated, then projected back to HIDDEN_DIM) so the K-loop signal doesn't vanish.
|
||||
- **Pro:** CfC stays load-bearing; the K-loop benefit (state amplification per `pearl_state_amplifies_short_horizon_into_long_horizon`) is preserved.
|
||||
- **Pro:** Minimal disruption to existing kernels — only attn_pool + multi_horizon_heads change shape; CfC + LN + Mamba2 + VSN unchanged.
|
||||
- **Con:** 5× attention compute + memory (manageable: K=6000, HIDDEN_DIM=128, FP32 → ~30 MB extra activation memory per batch, well within L40S budget).
|
||||
|
||||
### Path B — separate CfC per horizon
|
||||
- Each horizon gets its own context AND its own CfC instance with separate weights.
|
||||
- 5× the CfC compute + 5× the weight count.
|
||||
- **Pro:** Maximum specialisation.
|
||||
- **Con:** Capacity explosion (CfC weights are HIDDEN_DIM² = ~16k; 5× = 80k extra) + needs careful regularisation. High train-time cost.
|
||||
|
||||
### Path C — drop CfC, per-horizon contexts feed heads directly
|
||||
- Skip the K-loop entirely; rely on Mamba2-l2 + LN_b for the temporal modelling.
|
||||
- **Pro:** Simplest.
|
||||
- **Con:** Loses the K-loop / CfC benefit empirically validated in Phase 1+2+3. Likely regression.
|
||||
|
||||
**Default for v1: Path A.** Path B is a v2 refinement if Path A doesn't move the needle.
|
||||
|
||||
## 2. Kernel changes
|
||||
|
||||
Two new CUDA kernels replace `attention_pool.cu`:
|
||||
|
||||
`per_horizon_attention_pool_fwd`:
|
||||
```cuda
|
||||
grid_dim = (B, N_HORIZONS, 1)
|
||||
block_dim = (HIDDEN_DIM=128, 1, 1)
|
||||
```
|
||||
- Per-block: one (batch, horizon) pair handles HIDDEN_DIM threads → one context vector.
|
||||
- Block tree-reduce over HIDDEN_DIM for the per-k dot product (existing pattern from `attention_pool.cu`).
|
||||
- Per-horizon softmax over K (per-block within the (b, h) slice).
|
||||
- Final context[b, h, h_dim=tid] accumulation.
|
||||
- Per `feedback_no_atomicadd.md`: block tree-reduce only.
|
||||
|
||||
`per_horizon_attention_pool_bwd`:
|
||||
- Per-block (B, N_HORIZONS): single-writer discipline for `d_Q_h` per (h, dim).
|
||||
- `d_Q_h` is `[N_HORIZONS, HIDDEN_DIM]` shared across all batches. Per-block grad scratch is `[B, N_HORIZONS, HIDDEN_DIM]` then reduced via `reduce_axis0` (existing kernel from Phase B C1).
|
||||
- `d_LNb[b, k, d]` is per-batch indexed; each (b, h) block adds its contribution via plain `+=` (no race because batches are disjoint, horizons within a batch are summed serially within the block).
|
||||
|
||||
`multi_horizon_heads_fwd` extends signature:
|
||||
- Old: `(LNb_pooled[B, HIDDEN_DIM]) → probs[B, N_HORIZONS]`.
|
||||
- New: `(per_horizon_contexts[B, N_HORIZONS, HIDDEN_DIM], cfc_h_K[B, HIDDEN_DIM]) → probs[B, N_HORIZONS]`.
|
||||
- Each horizon head reads `context[b, h, :]` concat `cfc_h_K[b, :]` (residual) → projects to scalar logit.
|
||||
- Head weights: `head_w[N_HORIZONS, 2 * HIDDEN_DIM]` (was `[N_HORIZONS, HIDDEN_DIM]`).
|
||||
|
||||
`multi_horizon_heads_bwd` mirrors the above with the existing block-per-batch reduction pattern from Phase B C2.
|
||||
|
||||
## 3. Rust-side changes
|
||||
|
||||
`crates/ml-alpha/src/attention_pool.rs` (new — extract the existing inline binding from cfc/trunk.rs first if needed):
|
||||
- `pub struct AttentionPool { q_d: CudaSlice<f32> }` for single-query (legacy).
|
||||
- `pub struct PerHorizonAttentionPool { q_h_d: CudaSlice<f32> }` for the new path.
|
||||
- Both implement a common `pub trait AttentionPoolKernel { fn forward(...) -> Result<CudaSlice<f32>>; ... }` so the trainer can swap them via a config flag.
|
||||
|
||||
`crates/ml-alpha/src/trainer/perception.rs` config flag:
|
||||
```rust
|
||||
pub struct PerceptionTrainerConfig {
|
||||
// ... existing fields ...
|
||||
pub attention_pool_variant: AttentionPoolVariant,
|
||||
}
|
||||
|
||||
pub enum AttentionPoolVariant {
|
||||
SharedQuery, // current — Phase 1+2+3 default
|
||||
PerHorizonQuery, // new
|
||||
}
|
||||
```
|
||||
|
||||
Default stays `SharedQuery` to preserve all existing checkpoint compatibility. The new variant is opt-in via `--attention-pool per-horizon` CLI flag.
|
||||
|
||||
`CheckpointV1` bumps to `CheckpointV2`:
|
||||
- Add `attention_pool_variant: u8` discriminant.
|
||||
- Add `q_h: Option<Vec<f32>>` ([N_HORIZONS × HIDDEN_DIM] when PerHorizon).
|
||||
- Load path branches: V1 file → SharedQuery, V2 file → respect the discriminant.
|
||||
|
||||
## 4. Validation strategy
|
||||
|
||||
Three rings, in increasing rigour:
|
||||
|
||||
### Ring 1: kernel-level numgrad parity
|
||||
- Construct a tiny K=16, B=2, N_HORIZONS=5 fixture.
|
||||
- Run forward via the new kernels; capture probs.
|
||||
- Run backward; verify gradients match numerical FD-grad within `1e-4` tolerance (same approach as `Phase 2D.2 VSN backward numgrad check`).
|
||||
|
||||
### Ring 2: smoke training parity
|
||||
- Branch off the Phase 1+2+3 baseline checkpoint.
|
||||
- Train one epoch with `--attention-pool per-horizon`; assert no NaN, loss decreases, val AUC > chance (≥ 0.51).
|
||||
- This is a wiring check, not a quality check.
|
||||
|
||||
### Ring 3: A/B sweep (tracked as task #204)
|
||||
- 3-fold CV comparison:
|
||||
- Branch A: `--attention-pool shared-query` (existing baseline)
|
||||
- Branch B: `--attention-pool per-horizon`
|
||||
- 30 epochs each, identical seeds + folds + data.
|
||||
- Compare mean_auc per horizon + total. Per the falsifiable claim in §0, ship Branch B only if it lifts h6000 by ≥ +0.01 OR shifts the per-horizon distribution toward viable horizons (h1000, h300) with no net loss at h6000.
|
||||
|
||||
## 5. Risks + open questions
|
||||
|
||||
| Risk | Severity | Mitigation |
|
||||
|---|---|---|
|
||||
| 5× attention compute slows training meaningfully | Medium | Profile early. If wall regression > 20 %, reorder the attention loop to amortise (per-K row read once across all 5 horizons in the same block). |
|
||||
| Per-horizon Q_h params under-trained at short K | Low | First-observation bootstrap pattern (`pearl_first_observation_bootstrap.md`) doesn't apply to weight params; standard Xavier init + Phase 1d learning rate schedule. |
|
||||
| Auxiliary CfC h_K input creates spurious correlation that swamps per-horizon signal | Medium | Path A config flag `--head-aux-weight` defaults to 0.25 (heads consume 75% per-horizon + 25% CfC h_K). Tune if needed. |
|
||||
| Checkpoint compatibility breaks for existing trained models | Low | V1 → V2 migration is read-only (V1 files load as SharedQuery; new training writes V2 with discriminant). |
|
||||
| Bytecode VM in real-LOB sim doesn't know about per-horizon contexts | Out of scope | Real-LOB consumes only the inference `probs[N_HORIZONS]` output, which has the same shape in both variants. No backtest-side changes needed. |
|
||||
|
||||
### Open questions
|
||||
|
||||
1. **Q_h initialisation**: Xavier scaled to `1/sqrt(HIDDEN_DIM)` (same as existing Q) vs per-horizon-rescaled scaled to `1/sqrt(N_HORIZONS * HIDDEN_DIM)` (acknowledging the 5× param count). Default to the former; revisit if training-time variance is high.
|
||||
2. **Aux-weight schedule**: should `--head-aux-weight` decay from 1.0 → 0.0 over the first epoch (annealing) or stay constant? Constant for v1.
|
||||
3. **K-window length**: Path A assumes K=6000 (matches Phase 1+2+3). Per-horizon attention may benefit from shorter K for short horizons. Out of scope for v1 — fixed at K=6000.
|
||||
|
||||
## 6. Compliance with HARD memory rules
|
||||
|
||||
| Rule | How this spec complies |
|
||||
|---|---|
|
||||
| `feedback_no_atomicadd.md` | New per-horizon attention kernels use block tree-reduce, same as existing attention_pool.cu. |
|
||||
| `feedback_no_htod_htoh_only_mapped_pinned.md` | No new host↔device transfer paths. |
|
||||
| `feedback_no_nvrtc.md` | All kernels pre-compiled to cubins via existing build.rs in `crates/ml-alpha/`. |
|
||||
| `feedback_no_cpu_test_fallbacks.md` | Ring 1 numgrad parity is GPU-vs-GPU (FD grad on device). |
|
||||
| `feedback_no_partial_refactor.md` | Variant toggle is a config flag — old SharedQuery path remains fully functional; new PerHorizon path is opt-in. Checkpoint format bumps V1 → V2 with explicit migration. |
|
||||
| `pearl_state_amplifies_short_horizon_into_long_horizon.md` | Path A preserves the CfC K-loop entirely; per-horizon contexts are an ADDITIVE signal layered on top, not a replacement. |
|
||||
| `pearl_no_partial_refactor.md` | New kernels + trainer wiring + checkpoint bump land in a single branch (`ml-alpha-phase-c`); no half-shipped state. |
|
||||
|
||||
## 7. Implementation handoff
|
||||
|
||||
This spec is intentionally implementation-deferred. The decision to invest in per-horizon attention pool should be informed by:
|
||||
|
||||
1. **Time + GPU budget** for a 30-epoch × 3-fold A/B (≈ 6 × Phase 1+2+3 training time → ~3-6 hrs on L40S × 5 GPUs).
|
||||
2. **Whether the per-horizon cost-frontier sweep** (task #202 follow-up via `fxt-backtest sweep`) has identified additional viable horizons beyond h6000. Per `pearl_phase1d4_backtest_cost_edge_frontier.md` the cost edge for h6000 alone is already tight; broadening horizon viability is what makes per-horizon attention valuable.
|
||||
3. **The 3-fold variance baseline** — the `mean_auc 0.7749 ± 0.024` band means a +0.01 lift is within noise; a meaningful effect needs ≥ +0.024 or a per-horizon distributional shift that's qualitatively visible.
|
||||
|
||||
If those conditions are met, the next step is invoking `superpowers:writing-plans` against this spec to produce a multi-commit implementation plan (estimated 6-8 commits: 2 kernel pairs forward/backward, perception trainer wiring, checkpoint V2 bump, numgrad Ring 1, smoke Ring 2, A/B sweep Ring 3).
|
||||
|
||||
---
|
||||
|
||||
## Appendix: file inventory (for plan handoff)
|
||||
|
||||
New files:
|
||||
|
||||
```
|
||||
crates/ml-alpha/cuda/per_horizon_attention_pool.cu # NEW — fwd + bwd
|
||||
crates/ml-alpha/src/attention_pool.rs # NEW — variant trait + impls
|
||||
crates/ml-alpha/tests/per_horizon_attn_numgrad.rs # NEW — Ring 1 fixture
|
||||
```
|
||||
|
||||
Modified files:
|
||||
|
||||
```
|
||||
crates/ml-alpha/cuda/multi_horizon_heads.cu # extend forward signature to consume
|
||||
# per-horizon contexts + CfC h_K residual
|
||||
crates/ml-alpha/src/heads.rs # head_w shape [N_HORIZONS, 2 * HIDDEN_DIM]
|
||||
crates/ml-alpha/src/cfc/trunk.rs # bump CheckpointV1 → V2 with
|
||||
# attention_pool_variant + optional q_h
|
||||
crates/ml-alpha/src/trainer/perception.rs # AttentionPoolVariant config + wiring
|
||||
crates/ml-alpha/examples/alpha_train.rs # --attention-pool CLI flag
|
||||
infra/k8s/argo/alpha-perception-template.yaml # --attention-pool plumb-through
|
||||
scripts/argo-alpha-perception.sh # add --attention-pool flag handling
|
||||
```
|
||||
Reference in New Issue
Block a user