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