feat(temporal): streaming step() + KvCache (ADR-096 §3.2, #513)

The structural advantage that's the entire point of ADR-096: O(log T)
per new token via decode_step against an accumulated KvCache, vs
O(N²) recompute for dense MHA. This commit lands the API and proves
the numerical equivalence at the last position.

API:
- AetherTemporalHead::step(q_new, k_new, v_new, &mut cache)
  Single-token decode. Appends (k_new, v_new) to cache, runs
  decode_step(q_new) against the now-updated cache, returns the new
  position's output.
- AetherTemporalHead::make_cache(capacity)
  Convenience constructor — caller doesn't need to import
  ruvllm_sparse_attention to size a cache. Per ADR-096 §8.5 the
  natural lifetime is per-PoseTrack (re-ID) or per-session (online
  classification); when the track drops, drop the cache.
- KvCache re-exported at the crate root.

Contract:
- q_new/k_new/v_new must each have seq == 1. Multi-token q is the
  prefill path (forward), not decode_step.
- Cache lifetime is the caller's. The crate enforces shape via
  make_cache so callers can't mismatch kv_heads / head_dim / block_size.
- KvCache fill is the caller's problem. Upstream H2O heavy-hitter
  eviction is opt-in; this crate's wrapper doesn't pre-pick a policy.

Tests (18/18 total now passing):
- streaming_step_matches_forward_at_last_position — central claim:
  16-token sequence, append k/v one at a time via step(), compare
  the streamed last-token output to forward(full Q,K,V)[N-1].
  max_abs_err < 1e-3 (currently passes well under that bound for
  the 0.1-magnitude activations the test uses).
- step_rejects_multi_token_q — contract enforcement.
- make_cache_returns_kvcache_with_correct_shape — wiring smoke,
  confirms (capacity, kv_heads, dim, block_size) ordering is correct
  through the make_cache wrapper.

Test config uses MHA shape (q_heads == kv_heads) because the upstream
decode_step is wired to the MHA branch; the GQA decode path is on
upstream's roadmap and lands in a separate ADR-096 follow-up when it
does.

Co-Authored-By: claude-flow <ruv@ruv.net>
This commit is contained in:
ruv
2026-05-08 11:57:31 -04:00
parent 3a5fe5e0de
commit 49e57efcec
3 changed files with 216 additions and 4 deletions
+32 -3
View File
@@ -22,9 +22,9 @@ pub use weights::{
WEIGHT_BLOB_VERSION,
};
// Re-export the upstream Tensor3 so callers don't need a direct
// `ruvllm_sparse_attention` dep.
pub use ruvllm_sparse_attention::Tensor3;
// Re-export the upstream Tensor3 + KvCache so callers don't need a
// direct `ruvllm_sparse_attention` dep.
pub use ruvllm_sparse_attention::{KvCache, Tensor3};
/// Thin facade so callers can pick a backend by name.
///
@@ -62,4 +62,33 @@ impl AetherTemporalHead {
AetherTemporalHead::Dense => Err(TemporalError::DenseBackendNotImplemented),
}
}
/// Streaming decode (ADR-096 §3.2). Caller owns the `cache`; the
/// natural lifetime is per-tracked-person (one cache per
/// `PoseTrack`, dropped when the track evicts).
///
/// Returns the attention output for the single new token. Caller
/// is responsible for downstream pooling / classifier head.
pub fn step(
&self,
q_new: &Tensor3,
k_new: &Tensor3,
v_new: &Tensor3,
cache: &mut KvCache,
) -> Result<Tensor3, TemporalError> {
match self {
AetherTemporalHead::SparseGqa(h) => h.step(q_new, k_new, v_new, cache),
AetherTemporalHead::Dense => Err(TemporalError::DenseBackendNotImplemented),
}
}
/// Allocate a `KvCache` sized correctly for this head. Convenience
/// wrapper so AETHER's `pose_tracker.rs` doesn't need to import
/// the upstream crate.
pub fn make_cache(&self, capacity: usize) -> Result<KvCache, TemporalError> {
match self {
AetherTemporalHead::SparseGqa(h) => Ok(h.make_cache(capacity)),
AetherTemporalHead::Dense => Err(TemporalError::DenseBackendNotImplemented),
}
}
}
@@ -1,5 +1,5 @@
use ruvllm_sparse_attention::{
AttentionBackend, SparseAttentionConfig, SubquadraticSparseAttention, Tensor3,
AttentionBackend, KvCache, SparseAttentionConfig, SubquadraticSparseAttention, Tensor3,
};
use crate::{TemporalError, TemporalHeadConfig};
@@ -57,6 +57,50 @@ impl SparseGqaHead {
Ok(self.attn.forward_gqa(q, k, v)?)
}
}
/// Streaming decode for re-ID and online classification (ADR-096 §3.2).
///
/// Given one new token's q/k/v, append (k, v) to `cache` and return
/// the attention output for that one position against the full
/// accumulated history. Cost is O(log T) per step against a cache
/// of capacity T — the structural advantage over dense MHA's O(N²)
/// recompute that ADR-096 specifically calls out as the
/// dense-MHA-cannot-follow path.
///
/// Cache lifetime is owned by the caller. Per ADR-096 §8.5 the
/// natural place is one cache per `PoseTrack` (re-ID) or one cache
/// per active session (online classification). When the track is
/// dropped, drop the cache.
pub fn step(
&self,
q_new: &Tensor3,
k_new: &Tensor3,
v_new: &Tensor3,
cache: &mut KvCache,
) -> Result<Tensor3, TemporalError> {
if q_new.seq != 1 || k_new.seq != 1 || v_new.seq != 1 {
return Err(TemporalError::InvalidConfig(
"step() requires single-token q/k/v (seq == 1 each)",
));
}
// Append must succeed before decode_step sees the cache; if
// the cache fills, the caller is responsible for eviction or
// resetting per ADR-096 §3.2 (H2O heavy-hitter eviction is
// available upstream but kept opt-in).
cache.try_append(k_new, v_new)?;
Ok(self.attn.decode_step(q_new, cache)?)
}
/// Construct a KvCache sized for this head's shape. Convenience
/// so callers don't need to import the upstream crate directly.
pub fn make_cache(&self, capacity: usize) -> KvCache {
KvCache::new(
capacity,
self.cfg.kv_heads,
self.cfg.head_dim,
self.cfg.block_size,
)
}
}
/// Always treat token 0 as a global anchor — AETHER's contrastive