[PJRT:CPU] Use separate bounded thread pools for compilation, execution, and transfers.

- Introduce dedicated bounded thread pools in `PjRtCpuRawClient` for compilation (`compile_thread_pool_`) and execution (`execute_work_runner_`), separate from the general `async_work_runner_` used for buffer transfers and linearization/delinearization.
- Route `CompileInternal` to `compile_thread_pool_` and `CpuPjRtRawLoadedExecutable::Execute` to `execute_work_runner_`.

Separating execution, compilation, and transfer thread pools prevents multi-device collective launches from starving antecedent input buffer transfers on `async_work_runner_`, while keeping all three pools bounded so bursts of D2H/H2D buffer transfers remain throttled.

PiperOrigin-RevId: 987628676
This commit is contained in:
Junwhan Ahn
2026-09-25 14:07:44 -07:00
committed by TensorFlower Gardener
parent 0668dfb00c
commit 8bebdffc79
3 changed files with 20 additions and 7 deletions
-1
View File
@@ -248,7 +248,6 @@ cc_library(
"//xla/service/cpu:cpu_executable",
"//xla/service/cpu:cpu_executable_run_options",
"//xla/service/cpu:cpu_xfeed",
"//xla/service/cpu:executable_proto_cc",
"//xla/service/llvm_ir:llvm_command_line_options",
"//xla/stream_executor:device_address",
"//xla/tsl/concurrency:async_value",
+10 -6
View File
@@ -71,6 +71,7 @@ limitations under the License.
#include "xla/pjrt/cpu/cpu_async_execution_tracker.h"
#include "xla/pjrt/cpu/cpu_device_memory.h"
#include "xla/pjrt/cpu/cpu_event.h"
#include "xla/pjrt/cpu/execution_stream_event_map.h"
#include "xla/pjrt/cpu/raw_buffer.h"
#include "xla/pjrt/device_event.h"
#include "xla/pjrt/device_event_utils.h"
@@ -107,7 +108,6 @@ limitations under the License.
#include "xla/service/cpu/cpu_executable.h"
#include "xla/service/cpu/cpu_executable_run_options.h"
#include "xla/service/cpu/cpu_xfeed.h"
#include "xla/service/cpu/executable.pb.h"
#include "xla/service/device_assignment.h"
#include "xla/service/dump.h"
#include "xla/service/executable.h"
@@ -129,7 +129,6 @@ limitations under the License.
#include "xla/tsl/platform/threadpool.h"
#include "xla/util.h"
#include "xla/xla.pb.h"
#include "xla/xla_data.pb.h"
#include "tsl/platform/denormal.h"
#include "tsl/platform/fingerprint.h"
#include "tsl/platform/setround.h"
@@ -460,10 +459,15 @@ PjRtCpuRawClient::PjRtCpuRawClient(
eigen_intraop_pool_->NumThreads())
: std::unique_ptr<Eigen::ThreadPoolDevice>()),
custom_intraop_device_(intra_op_device),
compile_thread_pool_(std::make_unique<tsl::thread::ThreadPool>(
tsl::Env::Default(), GetThreadOptions(), "XLACompile", num_threads)),
execute_work_runner_(std::make_unique<ThreadPoolAsyncWorkRunner>(
tsl::Env::Default(), "XLAExecute", num_threads, GetThreadOptions())),
async_work_runner_(std::make_unique<ThreadPoolAsyncWorkRunner>(
tsl::Env::Default(), "XLAPjRtCpuClient", num_threads)) {}
tsl::Env::Default(), "XLAPjRtCpuClient", num_threads,
GetThreadOptions())) {}
PjRtCpuRawClient::~PjRtCpuRawClient() {}
PjRtCpuRawClient::~PjRtCpuRawClient() = default;
PjRtPluginAttributes GetDefaultCpuPluginAttributes() {
PjRtPluginAttributes attrs;
@@ -994,7 +998,7 @@ PjRtCpuRawClient::CompileInternal(
params.layout_canonicalization_callback =
std::move(layout_canonicalization_callback);
params.num_threads = num_threads;
params.compile_thread_pool = async_work_runner()->thread_pool();
params.compile_thread_pool = compile_thread_pool_.get();
params.aot_options = aot_options;
params.process_index = process_index;
params.collectives_exists = (collectives() != nullptr);
@@ -1697,7 +1701,7 @@ PjRtRawLoadedExecutable::RawExecuteResult CpuPjRtRawLoadedExecutable::Execute(
run_id_.ToInt(), std::move(ready_on_exit).Release());
PjRtDeviceEventSpan events_ref(input_deps);
xla::ExecuteWhenReady(
events_ref, raw_client->async_work_runner(),
events_ref, raw_client->execute_work_runner(),
[cpu_executable, buffer_alloc = std::move(buffer_alloc),
buffer_alloc_and_copy = std::move(buffer_alloc_and_copy),
execute_thunks = std::move(execute_thunks),
+10
View File
@@ -122,6 +122,14 @@ class PjRtCpuRawClient : public PjRtRawClient {
return async_work_runner_.get();
}
ThreadPoolAsyncWorkRunner* execute_work_runner() const {
return execute_work_runner_.get();
}
tsl::thread::ThreadPool* compile_thread_pool() const {
return compile_thread_pool_.get();
}
tsl::thread::ThreadPool* eigen_intraop_pool() const {
return eigen_intraop_pool_.get();
}
@@ -292,6 +300,8 @@ class PjRtCpuRawClient : public PjRtRawClient {
std::unique_ptr<tsl::thread::ThreadPool> eigen_intraop_pool_;
std::unique_ptr<Eigen::ThreadPoolDevice> eigen_intraop_device_;
const Eigen::ThreadPoolDevice* custom_intraop_device_ = nullptr;
std::unique_ptr<tsl::thread::ThreadPool> compile_thread_pool_;
std::unique_ptr<ThreadPoolAsyncWorkRunner> execute_work_runner_;
std::unique_ptr<ThreadPoolAsyncWorkRunner> async_work_runner_;
};