diff --git a/x/mlxrunner/cache/recurrent.go b/x/mlxrunner/cache/recurrent.go index d448533a5..c394161f9 100644 --- a/x/mlxrunner/cache/recurrent.go +++ b/x/mlxrunner/cache/recurrent.go @@ -40,7 +40,7 @@ type RecurrentCache struct { func (c *RecurrentCache) PrepareSnapshots(offsets []int) { c.snapshots.prepare(c.offset, offsets) // The current offset is a valid boundary right now, so capture it. - c.captureBoundary(c.offset) + c.captureBoundary(c.offset, c.convState, c.deltaState) } func (c *RecurrentCache) TakeSnapshots() []Snapshot { return c.snapshots.take() } @@ -62,32 +62,23 @@ func (c *RecurrentCache) SnapshotSplits(forwardLen int) []int { return splits } -func (c *RecurrentCache) captureBoundary(reached int) { - c.snapshots.captureReached(reached, func(int) Snapshot { return c.Snapshot(reached) }) -} - -// captureBoundaryState captures a scheduled interior offset from the -// per-boundary conv/delta states the kernel wrappers produced while segmenting, -// rather than from the cache's current (end-of-forward) state. -func (c *RecurrentCache) captureBoundaryState(reached int, conv, delta *mlx.Array) { +// captureBoundary snapshots the boundary state (conv, delta) at reached if +// reached is a scheduled offset — the live state at a forward boundary, or a +// kernel segment state at an interior split. +func (c *RecurrentCache) captureBoundary(reached int, conv, delta *mlx.Array) { c.snapshots.captureReached(reached, func(int) Snapshot { + // Nothing exists to page out yet; a zero-width capture is a nil entry. + if conv == nil && delta == nil { + return nil + } return newRecurrentSnapshot(conv, delta, reached) }) } -func (c *RecurrentCache) setState(old, v *mlx.Array, contiguous bool) *mlx.Array { - if v == nil || !v.Valid() { - return old - } - - if contiguous { - v = mlx.Contiguous(v, false) - } +func (c *RecurrentCache) setState(old, v *mlx.Array) *mlx.Array { v = v.Clone() - mlx.Pin(v) mlx.Unpin(old) - return v } @@ -123,8 +114,8 @@ func (c *RecurrentCache) Get(b *batch.Batch, dtype mlx.DType) *nn.RecurrentHisto return nn.NewRecurrentHistory(c.convState, c.deltaState) } - c.convState = c.setState(nil, mlx.Zeros(dtype, batch, c.convTail, c.convDim), false) - c.deltaState = c.setState(nil, mlx.Zeros(mlx.DTypeFloat32, batch, c.numVHeads, c.headVDim, c.headKDim), false) + c.convState = c.setState(nil, mlx.Zeros(dtype, batch, c.convTail, c.convDim)) + c.deltaState = c.setState(nil, mlx.Zeros(mlx.DTypeFloat32, batch, c.numVHeads, c.headVDim, c.headKDim)) return nn.NewRecurrentHistory(c.convState, c.deltaState) } @@ -154,19 +145,29 @@ func (c *RecurrentCache) Put(b *batch.Batch, convStates, deltaStates []*mlx.Arra // Leading entries are the interior split boundaries; capture each as a // snapshot at its scheduled offset. for i, s := range splits { - c.captureBoundaryState(start+s, convStates[i], deltaStates[i]) + c.captureBoundary(start+s, convStates[i], deltaStates[i]) } // The final entry is the forward-end state — the committed live state. last := len(convStates) - 1 - c.convState = c.setState(c.convState, convStates[last], true) - c.deltaState = c.setState(c.deltaState, deltaStates[last], false) + c.convState = c.setState(c.convState, convStates[last]) + c.deltaState = c.setState(c.deltaState, deltaStates[last]) c.offset += int(b.SeqQueryLens[0]) - c.captureBoundary(c.offset) + c.captureBoundary(c.offset, c.convState, c.deltaState) } func (c *RecurrentCache) State() []*mlx.Array { - return []*mlx.Array{c.convState, c.deltaState} + state := []*mlx.Array{c.convState, c.deltaState} + // Captured snapshots own compact copies still needing eval; fold them into + // the caller's batched State eval instead of an async eval per snapshot per + // layer. They drain at TakeSnapshots. + for _, s := range c.snapshots.captured { + if s != nil { + rs := s.(*recurrentSnapshot) + state = append(state, rs.convState, rs.deltaState) + } + } + return state } // recurrentSnapshot holds paged-out recurrent state. Self-contained — @@ -179,13 +180,13 @@ type recurrentSnapshot struct { func (s *recurrentSnapshot) Size() int { return s.convState.NumBytes() + s.deltaState.NumBytes() } func (s *recurrentSnapshot) Close() { mlx.Unpin(s.convState, s.deltaState) } -// SetMaterializeHook is a no-op: recurrent snapshots are always materialized -// at construction. +// SetMaterializeHook is a no-op: recurrent snapshots own their compact copy from +// construction. func (s *recurrentSnapshot) SetMaterializeHook(func(int)) {} // newRecurrentSnapshot clones and pins conv/delta into an owned snapshot at -// offset. Recurrent state is not position-sliceable, so a snapshot always owns -// a full copy. +// offset. It does not schedule the eval — capture-path snapshots ride the +// cache's State into the caller's batched eval. func newRecurrentSnapshot(conv, delta *mlx.Array, offset int) *recurrentSnapshot { snap := &recurrentSnapshot{ convState: conv.Clone(), @@ -202,7 +203,11 @@ func (c *RecurrentCache) Snapshot(fromOffset int) Snapshot { return nil } - return newRecurrentSnapshot(c.convState, c.deltaState, c.offset) + // Page-out snapshots go straight to the trie and never ride State, so + // schedule the eval here instead of through the batched State eval. + snap := newRecurrentSnapshot(c.convState, c.deltaState, c.offset) + mlx.AsyncEval(snap.convState, snap.deltaState) + return snap } func (c *RecurrentCache) Restore(snapshot Snapshot, target int) bool { @@ -221,8 +226,8 @@ func (c *RecurrentCache) Restore(snapshot Snapshot, target int) bool { return false } - c.convState = c.setState(c.convState, snap.convState, false) - c.deltaState = c.setState(c.deltaState, snap.deltaState, false) + c.convState = c.setState(c.convState, snap.convState) + c.deltaState = c.setState(c.deltaState, snap.deltaState) c.offset = snap.offset return true diff --git a/x/models/nn/recurrent.go b/x/models/nn/recurrent.go index 24fb32c76..57dedca57 100644 --- a/x/models/nn/recurrent.go +++ b/x/models/nn/recurrent.go @@ -146,7 +146,12 @@ func CausalConv1D(b *batch.Batch, input *mlx.Array, conv *Conv1d, convTail int, segs := segmentRanges(cfg.splits, L) states = make([]*mlx.Array, 0, len(segs)) for _, s := range segs { - states = append(states, convStateAt(concat, b.SeqQueryLens, convTail, s.end)) + st := convStateAt(concat, b.SeqQueryLens, convTail, s.end) + if L > 1 { + // Detach the small window from the forward-sized concat; at L==1 it's tiny. + st = mlx.Contiguous(st, false) + } + states = append(states, st) } return out, states }