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
2827from __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
0 commit comments