Skip to content

Commit 2c07d2b

Browse files
authored
Handle unused DeltaNet outputs in MLX (pytorch#23022)
Lower gated DeltaNet graphs when only the output tensor or final state remains live after dead-code elimination. Preserve persistent slots for live outputs. Extend the backend's existing op tests with output-only, state-only, and both-output cases for scan and custom kernels in FP32 and BF16. Check delegate and kernel selection and compare native outputs with the PyTorch reference. All 12 native cases passed. Without the fix, all eight single-output cases fail to lower while the four both-output controls still pass. MLX CI already runs this test module; no model download or Kev dependency is required. Authored with OpenAI Codex.
1 parent a65ba87 commit 2c07d2b

2 files changed

Lines changed: 84 additions & 46 deletions

File tree

‎backends/mlx/custom_kernel_ops/gated_delta_rule.py‎

Lines changed: 33 additions & 45 deletions
Original file line numberDiff line numberDiff line change
@@ -20,9 +20,8 @@
2020
getitem_1 = auto_func[1] # mutated state (BUFFER_MUTATION)
2121
return (getitem_1, getitem)
2222
23-
The pattern handler uses HEAD = getitem[1] (same as ETKVCacheUpdateHandler)
24-
because the partitioner needs the BUFFER_MUTATION node as a proper subgraph
25-
output. getitem[0] is left for the normal _getitem_handler to process.
23+
The handler preserves both outputs when present. A local state can leave only
24+
the output tensor or final state live after dead-code elimination.
2625
"""
2726

2827
from __future__ import annotations
@@ -122,8 +121,8 @@ class GatedDeltaRuleHandler(PatternHandler):
122121
"""
123122
Pattern for gated delta rule state mutation.
124123
125-
HEAD = getitem[1] (BUFFER_MUTATION — mutated state)
126-
BODY = [auto_func_node, getitem_0]
124+
HEAD = getitem[1] (mutated state), or getitem[0] when state is unused.
125+
BODY = auto_func_node and the other live getitem, if any.
127126
128127
Both getitem nodes are handled by this pattern to prevent
129128
_getitem_handler from calling slot_map on auto_func_node
@@ -136,7 +135,8 @@ def __init__(
136135
head: Node,
137136
body: List[Node],
138137
auto_func_node: Node,
139-
getitem_0: Node,
138+
getitem_0: Optional[Node],
139+
getitem_1: Optional[Node],
140140
q: Node,
141141
k: Node,
142142
v: Node,
@@ -147,6 +147,7 @@ def __init__(
147147
super().__init__(head, body)
148148
self.auto_func_node = auto_func_node
149149
self.getitem_0 = getitem_0
150+
self.getitem_1 = getitem_1
150151
self.q_node = q
151152
self.k_node = k
152153
self.v_node = v
@@ -169,12 +170,9 @@ def _is_auto_func_gated_delta_rule(node: Node) -> bool:
169170
def maybe_create(
170171
cls, ep: ExportedProgram, head: Node
171172
) -> Optional["GatedDeltaRuleHandler"]:
172-
"""
173-
Match HEAD = getitem[1] from auto_functionalized_v2(gated_delta_rule).
174-
"""
175173
if head.op != "call_function" or "getitem" not in str(head.target):
176174
return None
177-
if len(head.args) < 2 or head.args[1] != 1:
175+
if len(head.args) < 2 or head.args[1] not in (0, 1):
178176
return None
179177
if not isinstance(head.args[0], Node):
180178
return None
@@ -196,26 +194,26 @@ def maybe_create(
196194

197195
state = all_bases[0]
198196

199-
# Find getitem[0] (output tensor) among auto_func's users
200-
getitem_0 = None
197+
getitems = {}
201198
for user in auto_func_node.users:
202199
if (
203200
user.op == "call_function"
204201
and "getitem" in str(user.target)
205202
and len(user.args) >= 2
206-
and user.args[1] == 0
203+
and user.args[1] in (0, 1)
207204
):
208-
getitem_0 = user
209-
break
205+
getitems[user.args[1]] = user
210206

211-
if getitem_0 is None:
207+
if head is not getitems.get(1, getitems.get(0)):
212208
return None
213209

214210
return cls(
215211
head=head,
216-
body=[auto_func_node, getitem_0],
212+
body=[auto_func_node]
213+
+ [user for user in getitems.values() if user is not head],
217214
auto_func_node=auto_func_node,
218-
getitem_0=getitem_0,
215+
getitem_0=getitems.get(0),
216+
getitem_1=getitems.get(1),
219217
q=q,
220218
k=k,
221219
v=v,
@@ -224,6 +222,19 @@ def maybe_create(
224222
state=state,
225223
)
226224

225+
def _output_slots(self, P: MLXProgramBuilder) -> tuple[Slot, Slot]:
226+
out = (
227+
P.make_or_get_slot(self.getitem_0)
228+
if self.getitem_0 is not None
229+
else P.make_tmp_slot()[1]
230+
)
231+
carry = (
232+
P.make_or_get_slot(self.getitem_1)
233+
if self.getitem_1 is not None
234+
else P.make_tmp_slot()[1]
235+
)
236+
return out, carry
237+
227238
def __call__(self, P: MLXProgramBuilder, n: Node) -> Slot:
228239
assert n == self.head
229240

@@ -313,13 +324,7 @@ def _emit_metal_kernel(self, P: MLXProgramBuilder, n: Node) -> Slot:
313324

314325
# Output slot for y — use existing IO slot if getitem_0 is a graph output,
315326
# otherwise create a new temp slot.
316-
out = P.make_or_get_slot(self.getitem_0)
317-
318-
# Output slot for state_out (carry). This is node n's persistent output
319-
# (the mutated state), so it must be a node-owned slot — not a temp slot,
320-
# whose id is reclaimed on tmp_scope exit and would be read as dead by a
321-
# later node that consumes the mutated state (e.g. a second op call).
322-
carry = P.make_or_get_slot(n)
327+
out, carry = self._output_slots(P)
323328

324329
# Metal kernel source (non-vectorized, no mask variant from mlx-lm)
325330
source = """
@@ -440,11 +445,7 @@ def _emit_metal_kernel(self, P: MLXProgramBuilder, n: Node) -> Slot:
440445
)
441446
)
442447

443-
# HEAD is getitem[1] = mutated state → bind to carry
444-
# carry already registered as n's slot via make_or_get_slot(n) above.
445-
P.set_slot(self.getitem_0, out)
446-
447-
return carry
448+
return carry if self.getitem_1 is not None else out
448449

449450
def _emit_scan(self, P: MLXProgramBuilder, n: Node) -> Slot:
450451
"""Emit ScanNode decomposition of the gated delta recurrence."""
@@ -487,11 +488,7 @@ def _emit_scan(self, P: MLXProgramBuilder, n: Node) -> Slot:
487488
)
488489
q_slot, k_slot = q_exp, k_exp
489490

490-
# Carry needs a writable slot. This is node n's persistent output (the
491-
# mutated state), so it must be a node-owned slot — not a temp slot, whose
492-
# id is reclaimed on tmp_scope exit and would be read as dead by a later
493-
# node that consumes the mutated state (e.g. a second op call).
494-
carry = P.make_or_get_slot(n)
491+
out, carry = self._output_slots(P)
495492
P.emit(IdCopyNode(x=P.slot_to_tid(state_slot), out=P.slot_to_tid(carry)))
496493

497494
# Sliced temp slots for per-step inputs
@@ -501,9 +498,6 @@ def _emit_scan(self, P: MLXProgramBuilder, n: Node) -> Slot:
501498
_, g_s = P.make_tmp_slot()
502499
_, beta_s = P.make_tmp_slot()
503500

504-
# Output slot for the recurrence output.
505-
out = P.make_or_get_slot(self.getitem_0)
506-
507501
# Body temp slots
508502
_, t0 = P.make_tmp_slot()
509503
_, t1 = P.make_tmp_slot()
@@ -585,13 +579,7 @@ def _emit_scan(self, P: MLXProgramBuilder, n: Node) -> Slot:
585579
)
586580
)
587581

588-
# HEAD is getitem[1] = mutated state → bind to carry
589-
# carry already registered as n's slot via make_or_get_slot(n) above.
590-
591-
# Set getitem[0] slot → output tensor (for downstream computation)
592-
P.set_slot(self.getitem_0, out)
593-
594-
return carry
582+
return carry if self.getitem_1 is not None else out
595583

596584

597585
_registered = False

‎backends/mlx/custom_kernel_ops/test/test_gated_delta_rule.py‎

Lines changed: 51 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,3 @@
1-
#!/usr/bin/env python3
21
# Copyright (c) Meta Platforms, Inc. and affiliates.
32
# All rights reserved.
43
#
@@ -310,6 +309,56 @@ def create_inputs(self) -> Tuple[torch.Tensor, ...]:
310309
return (q, k, v, g, beta)
311310

312311

312+
class GatedDeltaRuleLocalStateModel(nn.Module):
313+
def __init__(self, outputs: str, use_custom_kernel: bool):
314+
super().__init__()
315+
self.outputs = outputs
316+
self.use_custom_kernel = use_custom_kernel
317+
318+
def forward(self, q, k, v, g, beta):
319+
state = q.new_zeros(q.shape[0], v.shape[2], v.shape[3], q.shape[3])
320+
out = torch.ops.mlx.gated_delta_rule(
321+
q, k, v, g, beta, state, use_custom_kernel=self.use_custom_kernel
322+
)
323+
if self.outputs == "state":
324+
return state
325+
if self.outputs == "output":
326+
return out
327+
return out, state
328+
329+
330+
class GatedDeltaRuleLocalStateTest(GatedDeltaRuleTest):
331+
def __init__(self, outputs: str, use_custom_kernel: bool, dtype: torch.dtype):
332+
super().__init__(
333+
batch_size=2,
334+
seq_len=5,
335+
num_heads=2,
336+
head_dim=32,
337+
value_dim=16,
338+
dtype=dtype,
339+
rtol=0.01 if dtype == torch.bfloat16 else 1e-4,
340+
atol=0.01 if dtype == torch.bfloat16 else 1e-4,
341+
use_custom_kernel=use_custom_kernel,
342+
)
343+
self.outputs = outputs
344+
self.name += f"_local_{outputs}"
345+
self.expected_node_counts = {
346+
"MetalKernelNode" if use_custom_kernel else "ScanNode": 1
347+
}
348+
349+
def create_model(self):
350+
return GatedDeltaRuleLocalStateModel(self.outputs, self.use_custom_kernel)
351+
352+
@classmethod
353+
def get_test_configs(cls):
354+
return [
355+
cls(outputs, kernel, dtype)
356+
for outputs in ("output", "state", "both")
357+
for kernel in (False, True)
358+
for dtype in (torch.float32, torch.bfloat16)
359+
]
360+
361+
313362
class GatedDeltaRuleDynamicSeqTest(OpTestCase):
314363
"""Test gated_delta_rule with dynamic seq_len.
315364
@@ -935,6 +984,7 @@ def create_inputs(self) -> Tuple[torch.Tensor, ...]:
935984

936985
configs = (
937986
GatedDeltaRuleTest.get_test_configs()
987+
+ GatedDeltaRuleLocalStateTest.get_test_configs()
938988
+ GatedDeltaRuleDynamicSeqTest.get_test_configs()
939989
+ GatedDeltaRuleGQATest.get_test_configs()
940990
+ GatedDeltaRuleFloatCastTest.get_test_configs()

0 commit comments

Comments
 (0)