From a1ae4db3a82863dee36b94de43f5881627a75911 Mon Sep 17 00:00:00 2001 From: Sohaib Iftikhar Date: Thu, 24 Sep 2026 10:26:13 -0700 Subject: [PATCH] [XLA:GPU]: Fix caching issue with kernels when unrolling while loops When while loops are unrolled the same thunk is called with `Record` multiple times for different loop iterations. Currently `CustomCallRecordState` stored a single sequence of `XLA_FFI_Command*`, so during update only the last recorded sequence of commands was passed to the FFI record API, overwriting the last iteration and leaving earlier iterations of the while loop with stale buffer pointers. Keying the state by the sink command pointer returned on `RecordCreate` (and passed back via `RecordUpdate::command`) so each unrolled iteration updates its own recorded command sequence. PiperOrigin-RevId: 987596254 --- .../backends/gpu/runtime/custom_call_thunk.cc | 35 ++++---- .../gpu/runtime/custom_call_thunk_test.cc | 90 +++++++++++++++++++ 2 files changed, 109 insertions(+), 16 deletions(-) diff --git a/third_party/xla/xla/backends/gpu/runtime/custom_call_thunk.cc b/third_party/xla/xla/backends/gpu/runtime/custom_call_thunk.cc index 2a2bc8f7f20..0ecd9b0d940 100644 --- a/third_party/xla/xla/backends/gpu/runtime/custom_call_thunk.cc +++ b/third_party/xla/xla/backends/gpu/runtime/custom_call_thunk.cc @@ -91,7 +91,9 @@ std::string GetSymbolName(const void* ptr) { } struct CustomCallRecordState : public CommandState { - std::vector commands; + absl::flat_hash_map> + commands; }; // A per-execution state that holds state for prepare and initialize stages. @@ -606,13 +608,14 @@ absl::StatusOr CustomCallThunk::Record( const bool is_record_create = std::holds_alternative(record_action); - const bool is_record_update = - std::holds_alternative(record_action); - if (is_record_update) { // Copy over commands from state to inline storage. - TF_RET_CHECK(state->commands.size() <= kMaxCommands) - << "Too many commands to fit in inline storage"; - std::copy(state->commands.begin(), state->commands.end(), commands_storage); - num_commands = state->commands.size(); + if (const auto* record_update = std::get_if(&record_action)) { + if (auto it = state->commands.find(record_update->command); + it != state->commands.end()) { + TF_RET_CHECK(it->second.size() <= kMaxCommands) + << "Too many commands to fit in inline storage"; + std::copy(it->second.begin(), it->second.end(), commands_storage); + num_commands = it->second.size(); + } } XLA_FFI_RecordAction action_to_pass = is_record_create @@ -679,19 +682,19 @@ absl::StatusOr CustomCallThunk::Record( command_buffer); } - // Save newly recorded commands to state if this is the Create action - // Must be done after returning from the FFI handler. - if (is_record_create) { - state->commands.assign(commands_storage, commands_storage + num_commands); - } - // Return the last command in the chain for dependency tracking. // If more than one command was recorded, and they are independent, a dummy // node must be added to the command graph by the FFI client so that XLA // can track a single dependency for the entire chain. if (num_commands > 0 && commands_storage[num_commands - 1] != nullptr) { - return reinterpret_cast( - commands_storage[num_commands - 1]); + const auto* sink_command = + reinterpret_cast( + commands_storage[num_commands - 1]); + if (is_record_create) { + state->commands[sink_command].assign(commands_storage, + commands_storage + num_commands); + } + return sink_command; } // No commands were recorded. return nullptr; diff --git a/third_party/xla/xla/backends/gpu/runtime/custom_call_thunk_test.cc b/third_party/xla/xla/backends/gpu/runtime/custom_call_thunk_test.cc index 63b6b6a03e9..459c6b35f04 100644 --- a/third_party/xla/xla/backends/gpu/runtime/custom_call_thunk_test.cc +++ b/third_party/xla/xla/backends/gpu/runtime/custom_call_thunk_test.cc @@ -1175,5 +1175,95 @@ TEST(CustomCallThunkTest, RecordCommandBufferFfiRecordWithEmptyCommand) { EXPECT_EQ(host_update, 90); } +TEST(CustomCallThunkTest, RecordCommandBufferMultipleRecordsInSameBuffer) { + ASSERT_OK_AND_ASSIGN(se::StreamExecutor * executor, GpuExecutor()); + if (executor->GetDeviceDescription().gpu_compute_capability().IsRocm()) { + GTEST_SKIP() << "AddI32 PTX kernel not supported on ROCm."; + } + + CustomCallThunk::OwnedHandlerBundle bundle; + bundle.execute = + ffi::Ffi::BindExecute().To([]() { return absl::OkStatus(); }); + bundle.record = ffi::Ffi::BindRecord() + .Ctx>() + .Arg() + .Arg() + .Ret() + .To(AddI32FfiHandler); + + ASSERT_OK_AND_ASSIGN(auto setup, FfiRecordTestSetup::Create( + executor, std::move(bundle), + "add_i32_ffi_unroll", {10, 20}, {0})); + + RecordTestAlloc alloc_iter1_create(executor); + ASSERT_OK_AND_ASSIGN( + auto slices_iter1_create, + AllocateAndCopy(*setup->stream, alloc_iter1_create, {1, 2}, {0})); + Thunk::ExecuteParams execute_params_iter1_create = + Thunk::ExecuteParams::Create( + ServiceExecutableRunOptions(), *alloc_iter1_create.buffer_allocations, + setup->stream.get(), setup->stream.get(), nullptr, nullptr, nullptr); + + ASSERT_OK_AND_ASSIGN(auto cb, executor->CreateCommandBuffer( + se::CommandBuffer::Mode::kPrimary)); + + // Record two iterations of the same CustomCallThunk into the same command + // buffer using the same CommandStateManager (as unrolled WhileThunk does). + ASSERT_OK_AND_ASSIGN( + const se::CommandBuffer::Command* cmd_0, + setup->thunk->Record(*setup->execute_params, *setup->record_params, + Command::RecordCreate{/*dependencies=*/{}}, + cb.get())); + ASSERT_OK_AND_ASSIGN( + const se::CommandBuffer::Command* cmd_1, + setup->thunk->Record(execute_params_iter1_create, *setup->record_params, + Command::RecordCreate{/*dependencies=*/{cmd_0}}, + cb.get())); + ASSERT_NE(cmd_0, cmd_1); + + ASSERT_OK(cb->Finalize()); + ASSERT_OK(cb->Submit(setup->stream.get())); + ASSERT_OK(setup->stream->BlockHostUntilDone()); + + // Update both recorded nodes with new buffer allocations. + ASSERT_OK(cb->Update()); + RecordTestAlloc alloc_iter0_update(executor); + ASSERT_OK_AND_ASSIGN( + auto slices_iter0_update, + AllocateAndCopy(*setup->stream, alloc_iter0_update, {40, 50}, {0})); + Thunk::ExecuteParams execute_params_iter0_update = + Thunk::ExecuteParams::Create( + ServiceExecutableRunOptions(), *alloc_iter0_update.buffer_allocations, + setup->stream.get(), setup->stream.get(), nullptr, nullptr, nullptr); + + RecordTestAlloc alloc_iter1_update(executor); + ASSERT_OK_AND_ASSIGN( + auto slices_iter1_update, + AllocateAndCopy(*setup->stream, alloc_iter1_update, {100, 200}, {0})); + Thunk::ExecuteParams execute_params_iter1_update = + Thunk::ExecuteParams::Create( + ServiceExecutableRunOptions(), *alloc_iter1_update.buffer_allocations, + setup->stream.get(), setup->stream.get(), nullptr, nullptr, nullptr); + + ASSERT_OK(setup->thunk->Record(execute_params_iter0_update, + *setup->record_params, + Command::RecordUpdate{cmd_0}, cb.get())); + ASSERT_OK(setup->thunk->Record(execute_params_iter1_update, + *setup->record_params, + Command::RecordUpdate{cmd_1}, cb.get())); + ASSERT_OK(cb->Finalize()); + ASSERT_OK(cb->Submit(setup->stream.get())); + ASSERT_OK(setup->stream->BlockHostUntilDone()); + + int32_t out_0 = 0; + int32_t out_1 = 0; + ASSERT_OK(setup->stream->Memcpy(&out_0, alloc_iter0_update.result_dev_ptrs[0], + sizeof(int32_t))); + ASSERT_OK(setup->stream->Memcpy(&out_1, alloc_iter1_update.result_dev_ptrs[0], + sizeof(int32_t))); + EXPECT_EQ(out_0, 90); + EXPECT_EQ(out_1, 300); +} + } // namespace } // namespace xla::gpu