Skip to content

Commit 54d4dba

Browse files
committed
fix failing python test
1 parent e42af0d commit 54d4dba

1 file changed

Lines changed: 9 additions & 7 deletions

File tree

python/src/coreai_models/export/mlir_ops.py

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)