@@ -296,9 +296,9 @@ def _replace_cache_update_autofuncs(
296296
297297 getitem_fetched = getitem_by_idx .get (0 ) # 4D fetched slice -> SDPA
298298 getitem_cache = getitem_by_idx .get (1 ) # 5D mutated cache -> handle
299- assert getitem_fetched is not None and getitem_cache is not None , (
300- f"{ autofunc_node .name } : expected getitem indices {{0, 1}} for "
301- f"(fetched slice, mutated cache); found { sorted (getitem_by_idx )} ."
299+ assert getitem_cache is not None , (
300+ f"{ autofunc_node .name } : no getitem at index 1 for the mutated cache; "
301+ f"the cache write would be dead-code-eliminated. Found { sorted (getitem_by_idx )} ."
302302 )
303303
304304 with graph .inserting_before (autofunc_node ):
@@ -361,11 +361,13 @@ def _replace_cache_update_autofuncs(
361361 args = (pre_squeeze , [0 ]),
362362 )
363363 # 4D meta comes from the fetched-slice getitem (what SDPA expects).
364- _copy_node_provenance (squeeze_op , getitem_fetched )
364+ if getitem_fetched is not None :
365+ _copy_node_provenance (squeeze_op , getitem_fetched )
365366
366- getitem_fetched .replace_all_uses_with (squeeze_op )
367- get_items .append (getitem_fetched )
368- get_item_replacements [getitem_fetched .name ] = squeeze_op
367+ if getitem_fetched is not None :
368+ getitem_fetched .replace_all_uses_with (squeeze_op )
369+ get_items .append (getitem_fetched )
370+ get_item_replacements [getitem_fetched .name ] = squeeze_op
369371
370372 getitem_cache .replace_all_uses_with (isu_node )
371373 get_items .append (getitem_cache )
0 commit comments