mirror of
https://github.com/tensorflow/tensorflow.git
synced 2026-09-28 05:13:36 +08:00
[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:
committed by
TensorFlower Gardener
parent
0668dfb00c
commit
8bebdffc79
Vendored
-1
@@ -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
@@ -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
@@ -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_;
|
||||
};
|
||||
|
||||
|
||||
Reference in New Issue
Block a user