mirror of
https://github.com/tensorflow/tensorflow.git
synced 2026-09-28 05:13:36 +08:00
[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
This commit is contained in:
committed by
TensorFlower Gardener
parent
eaedaed010
commit
a1ae4db3a8
+19
-16
@@ -91,7 +91,9 @@ std::string GetSymbolName(const void* ptr) {
|
||||
}
|
||||
|
||||
struct CustomCallRecordState : public CommandState {
|
||||
std::vector<const XLA_FFI_Command*> commands;
|
||||
absl::flat_hash_map<const se::CommandBuffer::Command*,
|
||||
absl::InlinedVector<const XLA_FFI_Command*, 1>>
|
||||
commands;
|
||||
};
|
||||
|
||||
// A per-execution state that holds state for prepare and initialize stages.
|
||||
@@ -606,13 +608,14 @@ absl::StatusOr<const se::CommandBuffer::Command*> CustomCallThunk::Record(
|
||||
|
||||
const bool is_record_create =
|
||||
std::holds_alternative<RecordCreate>(record_action);
|
||||
const bool is_record_update =
|
||||
std::holds_alternative<RecordUpdate>(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<RecordUpdate>(&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<const se::CommandBuffer::Command*> 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<const se::CommandBuffer::Command*>(
|
||||
commands_storage[num_commands - 1]);
|
||||
const auto* sink_command =
|
||||
reinterpret_cast<const se::CommandBuffer::Command*>(
|
||||
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;
|
||||
|
||||
@@ -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<ffi::Extension<ffi::RecordExtension>>()
|
||||
.Arg<ffi::AnyBuffer>()
|
||||
.Arg<ffi::AnyBuffer>()
|
||||
.Ret<ffi::AnyBuffer>()
|
||||
.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
|
||||
|
||||
Reference in New Issue
Block a user