From 9dcefa5272d2a2a828bbdb362435eca9bfc9615d Mon Sep 17 00:00:00 2001 From: Adam Durham Date: Tue, 31 Mar 2026 23:27:11 -0500 Subject: [PATCH] fix: break shared-buffer memory leak in GatedDeltaNet cache (#1077) Co-authored-by: Angelos Katharopoulos --- mlx_lm/models/kimi_linear.py | 2 +- mlx_lm/models/qwen3_5.py | 3 ++- mlx_lm/models/qwen3_next.py | 2 +- 3 files changed, 4 insertions(+), 3 deletions(-) diff --git a/mlx_lm/models/kimi_linear.py b/mlx_lm/models/kimi_linear.py index b2fa6d1..5730e9a 100644 --- a/mlx_lm/models/kimi_linear.py +++ b/mlx_lm/models/kimi_linear.py @@ -267,7 +267,7 @@ class ShortConv1d(nn.Module): positions = (ends[:, None] + mx.arange(n_keep))[..., None] new_state = mx.take_along_axis(conv_input, positions, axis=1) else: - new_state = conv_input[:, -n_keep:, :] + new_state = mx.contiguous(conv_input[:, -n_keep:, :]) return out, new_state diff --git a/mlx_lm/models/qwen3_5.py b/mlx_lm/models/qwen3_5.py index 16bb8b7..636a6ed 100644 --- a/mlx_lm/models/qwen3_5.py +++ b/mlx_lm/models/qwen3_5.py @@ -157,7 +157,8 @@ class GatedDeltaNet(nn.Module): qkv = mx.where(mask[..., None], qkv, 0) conv_input = mx.concatenate([conv_state, qkv], axis=1) if cache is not None: - cache[0] = conv_input[:, -(self.conv_kernel_size - 1) :] + n_keep = self.conv_kernel_size - 1 + cache[0] = mx.contiguous(conv_input[:, -n_keep:]) conv_out = nn.silu(self.conv1d(conv_input)) q, k, v = [ diff --git a/mlx_lm/models/qwen3_next.py b/mlx_lm/models/qwen3_next.py index 39128a4..d7e508b 100644 --- a/mlx_lm/models/qwen3_next.py +++ b/mlx_lm/models/qwen3_next.py @@ -266,7 +266,7 @@ class Qwen3NextGatedDeltaNet(nn.Module): positions = (ends[:, None] + mx.arange(n_keep))[..., None] cache[0] = mx.take_along_axis(conv_input, positions, axis=1) else: - cache[0] = conv_input[:, -n_keep:, :] + cache[0] = mx.contiguous(conv_input[:, -n_keep:, :]) conv_out = nn.silu(self.conv1d(conv_input))