Fix out-of-bounds crash in Stream affinity analysis for scf.while (#24723)

See issue https://github.com/iree-org/iree/issues/24722

The issue contains a deeper explanation of the issue and a repro case
(which is the same as the lit test here).

When an `scf.while` carries more values through its "before" region than
the op has results (i.e. some loop-carried values are not forwarded to a
result via `scf.condition`), the affinity analysis mapped an `scf.yield`
operand number directly onto `whileOp->getResult(operandNumber)`. Since
the "after" region yield has one operand per init operand (`N`) while
the op may have fewer results (`M`), this lead to `getResult()` being
out of bounds.

Two independent sites made this naive mapping:

- **`Affinity.cpp`**: in the `scf.yield` → `scf.while` branch of
`ValueConsumerAffinityPVS::updateFromUse`, drop the
`whileOp->getResult(operandNumber)` propagation and keep only the
before-region-argument propagation. (Results are already handled by the
`scf.condition` case.)

- **`Explorer.cpp`**: guard the `ReturnLike` result-mapping fallback in
`walkTransitiveUses` with
`!isa<RegionBranchTerminatorOpInterface>(ownerOp)` (region-branch
terminators like `scf.yield` are already handled correctly above) plus a
`use.getOperandNumber() < parent->getNumResults()` bounds check.

Added regression test `@scf_while_extra_loop_carried`, which exercises
the `N > M` shape (two loop-carried values, one result) that previously
asserted.

Note that Claude made the fix suggestions and, while they look okay to
me, I'm not at all familiar with this code.

---------

Signed-off-by: Paul Stark <paul.stark@cdprojektred.com>
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
Co-authored-by: Artem Gindinson <gindinson@roofline.ai>
diff --git a/compiler/src/iree/compiler/Dialect/Stream/Analysis/Affinity.cpp b/compiler/src/iree/compiler/Dialect/Stream/Analysis/Affinity.cpp
index a1ea379..5983417 100644
--- a/compiler/src/iree/compiler/Dialect/Stream/Analysis/Affinity.cpp
+++ b/compiler/src/iree/compiler/Dialect/Stream/Analysis/Affinity.cpp
@@ -543,17 +543,15 @@
           return TraversalResult::COMPLETE;
         } else if (auto whileOp =
                        dyn_cast<mlir::scf::WhileOp>(op->getParentOp())) {
+          // The yield terminates the "after" region and forwards its operands
+          // back to the "before" region block arguments, matching the while's
+          // init operands, *not* the while results - those are produced by the
+          // scf.condition op and handled in the dedicated case.
           auto value = Position::forValue(
               whileOp.getBefore().getArgument(operand.getOperandNumber()));
           auto &valueUsage = solver.getElementFor<ValueConsumerAffinityPVS>(
               *this, value, DFX::Resolution::REQUIRED);
           newState ^= valueUsage.getState();
-          auto &parentUsage = solver.getElementFor<ValueConsumerAffinityPVS>(
-              *this,
-              Position::forValue(
-                  whileOp->getResult(operand.getOperandNumber())),
-              DFX::Resolution::REQUIRED);
-          newState ^= parentUsage.getState();
           return TraversalResult::COMPLETE;
         } else if (auto forOp = dyn_cast<mlir::scf::ForOp>(op->getParentOp())) {
           auto value = Position::forValue(
diff --git a/compiler/src/iree/compiler/Dialect/Stream/Transforms/test/annotate_affinities.mlir b/compiler/src/iree/compiler/Dialect/Stream/Transforms/test/annotate_affinities.mlir
index f73cfff..40ba703 100644
--- a/compiler/src/iree/compiler/Dialect/Stream/Transforms/test/annotate_affinities.mlir
+++ b/compiler/src/iree/compiler/Dialect/Stream/Transforms/test/annotate_affinities.mlir
@@ -1567,6 +1567,116 @@
 
 // -----
 
+// Tests an scf.while that carries more values through its "before" region than
+// the op has results: the "after" region yield feeds back two values (matching
+// the two init operands) while scf.condition only forwards one to the single
+// result.
+
+// CHECK-LABEL: @scf_while_extra_loop_carried
+util.func private @scf_while_extra_loop_carried() -> tensor<1xi32> {
+  %c0 = arith.constant 0 : index
+  %c2_i32 = arith.constant 2 : i32
+  // CHECK: flow.tensor.constant
+  // CHECK-SAME{LITERAL}: stream.affinities.results = [[#hal.device.promise<@dev_a>]]
+  %cst_a = flow.tensor.constant {stream.affinity = #hal.device.promise<@dev_a>} dense<123> : tensor<1xi32>
+  // CHECK: flow.tensor.constant
+  // CHECK-SAME{LITERAL}: stream.affinities.results = [[#hal.device.promise<@dev_b>]]
+  %cst_b = flow.tensor.constant {stream.affinity = #hal.device.promise<@dev_b>} dense<456> : tensor<1xi32>
+  // CHECK: scf.while
+  %while = scf.while(%arg0 = %cst_a, %arg1 = %cst_b) : (tensor<1xi32>, tensor<1xi32>) -> tensor<1xi32> {
+    // CHECK: flow.tensor.load
+    // CHECK-SAME{LITERAL}: stream.affinities.operands = [[#hal.device.promise<@dev_a>]]
+    %cond_i32 = flow.tensor.load %arg0[%c0] : tensor<1xi32>
+    %cond = arith.cmpi slt, %cond_i32, %c2_i32 : i32
+    // Only %arg0 is forwarded to the single result.
+    // CHECK: scf.condition
+    // CHECK-SAME{LITERAL}: stream.affinities.operands = [[#hal.device.promise<@dev_a>]]
+    scf.condition(%cond) %arg0 : tensor<1xi32>
+  } do {
+  ^bb0(%arg0: tensor<1xi32>):
+    // CHECK: flow.dispatch @dispatch
+    // CHECK-SAME{LITERAL}: stream.affinities.results = [[#hal.device.promise<@dev_a>]]
+    %t = flow.dispatch @dispatch(%arg0) : (tensor<1xi32>) -> tensor<1xi32>
+    // Yields two values back into the "before" region; operand #1 has no
+    // matching while result.
+    // CHECK: scf.yield
+    // CHECK-SAME{LITERAL}: stream.affinities.operands = [[#hal.device.promise<@dev_a>], [#hal.device.promise<@dev_a>]]
+    scf.yield %t, %arg0 : tensor<1xi32>, tensor<1xi32>
+  // CHECK: } attributes {
+  // CHECK-SAME{LITERAL}: stream.affinities.operands = [[#hal.device.promise<@dev_a>], [#hal.device.promise<@dev_b>]]
+  // CHECK-SAME{LITERAL}: stream.affinities.results = [[#hal.device.promise<@dev_a>]]
+  }
+  // CHECK: flow.tensor.transfer
+  // CHECK-SAME{LITERAL}: stream.affinities.operands = [[#hal.device.promise<@dev_a>]]
+  // CHECK-SAME{LITERAL}: stream.affinities.results = [[#hal.device.promise<@dev_c>]]
+  %while_c = flow.tensor.transfer %while : tensor<1xi32> to #hal.device.promise<@dev_c>
+  // CHECK: util.return
+  // CHECK-SAME{LITERAL}: stream.affinities.operands = [[#hal.device.promise<@dev_c>]]
+  util.return %while_c : tensor<1xi32>
+}
+
+// -----
+
+// Companion to @scf_while_extra_loop_carried where the value forwarded to the
+// single result is the *second* loop-carried value (before-region argument
+// index 1), and its init constant is left unpinned. This checks that the while
+// consumer's affinity (dev_c) propagates back to the correct before-region
+// argument by index and reaches a value with no explicit affinity of its own.
+
+// CHECK-LABEL: @scf_while_return_second_loop_carried
+util.func private @scf_while_return_second_loop_carried() -> tensor<1xi32> {
+  %c0 = arith.constant 0 : index
+  %c2_i32 = arith.constant 2 : i32
+  // Pinned; drives the loop condition but is not returned.
+  // CHECK: flow.tensor.constant
+  // CHECK-SAME{LITERAL}: stream.affinities.results = [[#hal.device.promise<@dev_a>]]
+  %cst_a = flow.tensor.constant {stream.affinity = #hal.device.promise<@dev_a>} dense<123> : tensor<1xi32>
+  // Not pinned: its affinity is derived from the while consumer (dev_c).
+  // CHECK: flow.tensor.constant
+  // CHECK-SAME{LITERAL}: stream.affinities.results = [[#hal.device.promise<@dev_c>]]
+  %cst_b = flow.tensor.constant dense<456> : tensor<1xi32>
+  // CHECK: scf.while
+  %while = scf.while(%arg0 = %cst_a, %arg1 = %cst_b) : (tensor<1xi32>, tensor<1xi32>) -> tensor<1xi32> {
+    // CHECK: flow.tensor.load
+    // CHECK-SAME{LITERAL}: stream.affinities.operands = [[#hal.device.promise<@dev_a>]]
+    %cond_i32 = flow.tensor.load %arg0[%c0] : tensor<1xi32>
+    %cond = arith.cmpi slt, %cond_i32, %c2_i32 : i32
+    // Only %arg1 is forwarded to the single result.
+    // CHECK: scf.condition
+    // CHECK-SAME{LITERAL}: stream.affinities.operands = [[#hal.device.promise<@dev_c>]]
+    scf.condition(%cond) %arg1 : tensor<1xi32>
+  } do {
+  ^bb0(%arg0: tensor<1xi32>):
+    // CHECK: flow.dispatch @dispatch
+    // CHECK-SAME{LITERAL}: stream.affinities.results = [[#hal.device.promise<@dev_c>]]
+    %t = flow.dispatch @dispatch(%arg0) : (tensor<1xi32>) -> tensor<1xi32>
+    // Yields back into the "before" region: before-arg #0 <- %cst_a (dev_a),
+    // before-arg #1 <- %t (dev_c).
+    // CHECK: scf.yield
+    // CHECK-SAME{LITERAL}: stream.affinities.operands = [[#hal.device.promise<@dev_a>], [#hal.device.promise<@dev_c>]]
+    scf.yield %cst_a, %t : tensor<1xi32>, tensor<1xi32>
+  // CHECK: } attributes {
+  // CHECK-SAME{LITERAL}: stream.affinities.operands = [[#hal.device.promise<@dev_a>], [#hal.device.promise<@dev_c>]]
+  // The while result's producer affinity is currently over-approximated as the
+  // union of both loop-carried affinities ([dev_a, dev_c]) instead of just
+  // [dev_c] (the single result only forwards %arg1). Re-enable this check once
+  // that producer-side over-approximation is fixed (https://github.com/iree-org/iree/issues/24787).
+  // COM: CHECK-SAME{LITERAL}: stream.affinities.results = [[#hal.device.promise<@dev_c>]]
+  }
+  // CHECK: flow.tensor.transfer
+  // CHECK-SAME{LITERAL}: stream.affinities.results = [[#hal.device.promise<@dev_c>]]
+  // The transfer consumes the while result, so its operand affinity inherits
+  // the same over-approximated union. Re-enable this check once the
+  // producer-side over-approximation is fixed (https://github.com/iree-org/iree/issues/24787).
+  // COM: CHECK-SAME{LITERAL}: stream.affinities.operands = [[#hal.device.promise<@dev_c>]]
+  %while_c = flow.tensor.transfer %while : tensor<1xi32> to #hal.device.promise<@dev_c>
+  // CHECK: util.return
+  // CHECK-SAME{LITERAL}: stream.affinities.operands = [[#hal.device.promise<@dev_c>]]
+  util.return %while_c : tensor<1xi32>
+}
+
+// -----
+
 // Tests a realistic program with ABI ops.
 
 // CHECK-LABEL: @simple_program
diff --git a/compiler/src/iree/compiler/Dialect/Util/Analysis/Explorer.cpp b/compiler/src/iree/compiler/Dialect/Util/Analysis/Explorer.cpp
index 9c52400..9b921c2 100644
--- a/compiler/src/iree/compiler/Dialect/Util/Analysis/Explorer.cpp
+++ b/compiler/src/iree/compiler/Dialect/Util/Analysis/Explorer.cpp
@@ -1277,12 +1277,10 @@
         result |= traverseReturnOp(ownerOp, use.getOperandNumber());
       }
 
-      if (ownerOp->hasTrait<OpTrait::ReturnLike>() &&
-          !isa<CallableOpInterface>(ownerOp->getParentOp())) {
-        auto parent = ownerOp->getParentOp();
-        auto result = parent->getResult(use.getOperandNumber());
-        worklist.insert(result);
-      }
+      // Note that we are not explicitly handling the case of a return op
+      // without a callable parent here since the above
+      // `traverseRegionBranchOp` already handled it (`ReturnLike` implies
+      // `RegionBranchTerminatorOpInterface`).
 
       // Step across global stores and into all of the loads across the program.
       if (auto storeOp =