cache: stop recurrent conv state from pinning the forward buffer
Keeping the recurrent conv state small was handled unevenly: the committed live state was recopied on every commit — wasted work on single-token decode, where the window is already tiny — while boundary states captured as snapshots could still be plain slices of the forward-sized convolution buffer. A cached slice pins that whole buffer even though the trie's eviction accounting only counts the slice's bytes, so recurrent cache memory piled up across requests and eviction could never reclaim it. Compact each boundary state to its real size once, where it is produced in the conv wrapper, so live state and snapshots own only their own bytes and eviction sees the true cost. Single-token decode leaves the already-tiny window as a slice. Fixes #16698
This commit is contained in:
Vendored
+38
-33
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user