- go →auto
This is what leaves JAX and crosses the seam, the form every chapter at the waist works on.
%0 = stablehlo.dot_general %arg0, %arg1, contracting_dims = [1] x [0], precision = [DEFAULT, DEFAULT] : (tensor<16x32xf32>, tensor<32x32xf32>) -> tensor<16x32xf32>
%1 = stablehlo.dot_general %arg0, %arg2, contracting_dims = [1] x [0], precision = [DEFAULT, DEFAULT] : (tensor<16x32xf32>, tensor<32x32xf32>) -> tensor<16x32xf32>
%2 = stablehlo.dot_general %arg0, %arg3, contracting_dims = [1] x [0], precision = [DEFAULT, DEFAULT] : (tensor<16x32xf32>, tensor<32x32xf32>) -> tensor<16x32xf32>
%3 = stablehlo.transpose %1, dims = [1, 0] : (tensor<16x32xf32>) -> tensor<32x16xf32>
%4 = stablehlo.dot_general %0, %3, contracting_dims = [1] x [0], precision = [DEFAULT, DEFAULT] : (tensor<16x32xf32>, tensor<32x16xf32>) -> tensor<16x16xf32>
%cst = stablehlo.constant dense<3.200000e+01> : tensor<f32>
%5 = stablehlo.sqrt %cst : tensor<f32>
%6 = stablehlo.broadcast_in_dim %5, dims = [] : (tensor<f32>) -> tensor<16x16xf32>
%7 = stablehlo.divide %4, %6 : tensor<16x16xf32> How to read this chapter
An abstract class is a list of promises with no bodies, so this chapter's body is code: three guided walks below step through the real implementations line by line. The first walks the CPU client, the in-process implementation you can single-step on a laptop, from CompileAndLoad down to the thunk run, with asides for where the GPU client does the same job with streams and NCCL. The second walks the PJRT-backed IFRT adapter keeping chapter 11's promises by delegation, and ends on the proxy keeping the very same promise over a wire. The third walks torch_xla's PjRtComputationClient, which sits above the seam rather than under it, one round trip from a parameter leaving host memory to a literal coming back. The short sections here hold only what the walked excerpts cannot show, and LAB·X5 at the bottom is the same interface implemented from scratch: a mock GPU you build in plain C. The XLAThe compiler: brilliant at fusing along dataflow edges, structurally unable to change your algorithm. That gap is why kernels exist.taught in /l/xla → code was read at openxla/xla commit 881f236 on 2026-08-10, and the torch_xla code at pytorch/xla commit 41398bf.
The cast, so every name has a home. PjRtCpuClient (xla/pjrt/cpu/cpu_client.h) and StreamExecutorGpuClient (xla/pjrt/gpu/se_gpu_pjrt_client.h, built on PjRtStreamExecutorClient, now under xla/pjrt/se/) implement PJRT in-process over a shared CommonPjRtClient base. The C API and its client-side wrapper sit in xla/pjrt/c/ and xla/pjrt/c_api_client/. The PJRT-backed IFRT adapter is xla/python/pjrt_ifrt/, and the proxy pair is xla/python/ifrt_proxy/.
The contract around the walked code
Read Execute's signature slowly, because every noun chapter 1 introduced is in it. Arguments arrive as a span of vectors of raw PjRtBuffer pointers, one inner vector per partition, exactly the nested list chapter 1 described, and the adapter walk's transpose step shows who assembles it. Results come back as vectors of unique_ptr, and the asymmetry is the ownership story: the caller lends its input buffers and owns every output outright. The optional futures are how a caller asks to be told, per device, when execution really completes.
The definition event the buffer walk-steps introduce reaches further than those excerpts show. Everything on a buffer queues behind it: ToLiteral returns a Future<> because the data may not exist yet, Delete() drops the caller's claim at once but frees memory only after every enqueued reader finishes, and donation rides the same bookkeeping in reverse, an input surrendered to Execute becoming eligible output storage, with ExecuteOptions' non_donatable_input_indices as the opt-out for arguments a caller wants to keep.
virtual absl::StatusOr<std::vector<std::vector<std::unique_ptr<PjRtBuffer>>>>
Execute(absl::Span<const std::vector<PjRtBuffer*>> argument_handles,
const ExecuteOptions& options,
std::optional<std::vector<Future<>>>& returned_futures) const = 0; The ABI crossing
The walk ends at GetPjrtApi returning a struct of function pointers; this is what happens on the other side of it. PJRT_Api_Version carries a major and a minor (0 and 114 at this reading), the plugin reports the pair it compiled against, and every args struct's struct_size lets a newer caller and an older plugin read only as much of each other as they both know. Optional capability chains off PJRT_Extension_Base structs so extensions never touch the core ABI.
And the frontend never calls the table directly. PjRtCApiClient (xla/pjrt/c_api_client/pjrt_c_api_client.h) wraps the function pointers back into the same C++ PjRtClient interface, while XLAThe compiler: brilliant at fusing along dataflow edges, structurally unable to change your algorithm. That gap is why kernels exist.taught in /l/xla →'s own plugins implement the table by wrapping a real C++ client, which the excerpt below catches in the act: args->client->client->CompileAndLoad(...), the outer client a C struct, the inner one a C++ object. A program crossing the plugin boundary meets the C++ interface twice, with a function table and two protobufs in between and nothing else.
PJRT_ASSIGN_OR_RETURN(
std::unique_ptr<xla::PjRtLoadedExecutable> executable,
std::visit(absl::Overload{
[args, &options](xla::MaybeOwningMlirModule module) {
return args->client->client->CompileAndLoad(
std::move(module), options);
},
[args, &options](xla::XlaComputation program) {
return args->client->client->CompileAndLoad(program,
options);
},
},
std::move(module_or_hlo))); The proxy, and the third route
The proxy client (xla/python/ifrt_proxy/client/client.cc) implements the same ifrt::Client functions the adapter walk stepped through, a second way: each one serializes its arguments into a protobuf request, one message type per interface function, and the server's IfrtBackend (xla/python/ifrt_proxy/server/ifrt_backend.cc) is a switch over every request case, replaying calls onto whichever in-process client it wraps. Read the two files side by side and the interface appears a third time, as a protocol: every promise in client.h has a proto twin.
Which is the general lesson this chapter has been circling. An interface can be implemented by doing the work, the CPU client; by delegating to something that does, the adapter and the C API wrapper; or by shipping the call to a process that does, the proxy. XLAThe compiler: brilliant at fusing along dataflow edges, structurally unable to change your algorithm. That gap is why kernels exist.taught in /l/xla →'s tree holds all three in the open. The fourth, keeping the promises with machinery that is not XLA's at all, is chapter 13 standing behind the same seams, and nothing in the code you just walked would notice.
Guided walks
1absl::StatusOr<std::unique_ptr<PjRtLoadedExecutable>>2PjRtCpuClient::CompileAndLoad(const XlaComputation& computation,3 CompileOptions options) {4 ABSL_ASSIGN_OR_RETURN(auto results,5 CompileAndAssignDevices(computation, std::move(options)));6 return LoadInternal(std::move(results.first), std::move(results.second));7}8absl::StatusOr<std::pair<std::unique_ptr<PjRtCpuExecutable>,9 std::shared_ptr<DeviceAssignment>>>10PjRtCpuClient::CompileAndAssignDevices(MaybeOwningMlirModule module,11 CompileOptions options) {12 ABSL_ASSIGN_OR_RETURN(MlirCompilationSetup setup,13 SetupMlirCompilation(module, options, *topology_));14 // ... layout resolution trimmed ...15 return CompileInternal(setup.computation, setup.argument_layout_pointers,16 setup.layout_callback, options,17 /*aot_options=*/nullptr);18}19static absl::StatusOr<std::unique_ptr<xla::Executable>> JitCompile(20 std::unique_ptr<HloModule> hlo_module,21 const ExecutableBuildOptions& build_options,22 const ExecutionOptions& execution_options,23 const xla::Compiler::CompileOptions& compile_options) {24 // ... LLVM option plumbing trimmed ...25 cpu::CpuCompiler compiler;26 if (!build_options.run_backend_only()) {27 ABSL_ASSIGN_OR_RETURN(hlo_module, compiler.RunHloPasses(std::move(hlo_module),28 /*stream_exec=*/nullptr,29 compile_options));30 }31 return compiler.RunBackend(std::move(hlo_module), /*stream_exec=*/nullptr,32 compile_options);33}34absl::StatusOr<std::unique_ptr<PjRtExecutable>> PjRtCpuClient::Compile(35 const XlaComputation& computation, CompileOptions options) {36 ABSL_ASSIGN_OR_RETURN(auto results,37 CompileAndAssignDevices(computation, std::move(options)));38 return std::move(results.first);39}40CommonPjRtClient::BufferFromHostBuffer(41 const void* data, PrimitiveType type, absl::Span<int64_t const> dims,42 std::optional<absl::Span<int64_t const>> byte_strides,43 HostBufferSemantics host_buffer_semantics,44 absl::AnyInvocable<void() &&> on_done_with_host_buffer,45 PjRtMemorySpace* memory_space, const Layout* device_layout) {46 // ... shape validation trimmed ...47 if (host_buffer_semantics ==48 PjRtClient::HostBufferSemantics::kImmutableZeroCopy ||49 host_buffer_semantics ==50 PjRtClient::HostBufferSemantics::kMutableZeroCopy) {51 if (BufferFromHostBufferSupportsZeroCopy(data, type, dims, byte_strides,52 *shared_device_shape, memory_space,53 device_layout)) {54 ABSL_ASSIGN_OR_RETURN(55 auto raw_buffer,56 ImportForeignMemory(57 const_cast<void*>(data),58 std::move(on_done_with_host_buffer), on_device_bytes_count,59 memory_space, /* ... */));60 // ... wraps raw_buffer into the returned PjRtBuffer ...61 }62 }63 ABSL_ASSIGN_OR_RETURN(auto raw_buffer,64 AllocateRawBuffer(memory_space, on_device_bytes_count,65 /*retry_on_oom=*/true,66 /*allocate_after=*/{}));67 ABSL_ASSIGN_OR_RETURN(auto definition_event,68 LinearizeHostBufferInto(data, type, dims, byte_strides,69 host_buffer_semantics,70 std::move(on_done_with_host_buffer),71 *shared_device_shape, raw_buffer));72 ABSL_ASSIGN_OR_RETURN(std::unique_ptr<PjRtBuffer> output_buffer,73 DefineBuffer(shared_device_shape, memory_space, raw_buffer,74 {std::move(definition_event)}));75 return output_buffer;76}77 cpu::BufferAllocations allocations(buffer_device_mem);7879 ABSL_ASSIGN_OR_RETURN(cpu::Thunk::CollectiveExecuteParams collective_params,80 cpu::Thunk::CollectiveExecuteParams::Create(&run_options));81 // ... custom-call params and task-runner setup trimmed ...82 cpu::Thunk::ExecuteParams execute_params = {83 cpu_executable->function_library(),84 &allocations,85 /* ... */86 };87 auto thunks_execute_event =88 cpu_executable->thunks().Execute(execute_params);89 tsl::BlockUntilReady(thunks_execute_event);
01/07The loaded flavor of compilation, top of the chain. Two calls carry the whole promise: CompileAndAssignDevices produces the compiled artifact plus a device assignment, and LoadInternal binds them together into the executable chapter 1 met, ready to run. Every real client has a chain like this; the GPU's StreamExecutorGpuClient ends in the same two-step, with its artifact aimed at a different chip.
1absl::StatusOr<ArrayRef> PjRtClient::MakeArrayFromHostBuffer(2 const void* data, DType dtype, Shape shape,3 std::optional<absl::Span<const int64_t>> byte_strides, ShardingRef sharding,4 LayoutRef layout, Client::HostBufferSemantics semantics,5 std::function<void()> on_done_with_host_buffer) {6 if (dtype.kind() == DType::kString) {7 return MakeStringArrayFromHostBuffer(this, data, dtype, shape, byte_strides,8 sharding, semantics,9 on_done_with_host_buffer);10 }11 if (!isa<const SingleDeviceSharding>(sharding.get()) &&12 !sharding->IsFullyReplicated()) {13 return InvalidArgument(14 "Only SingleDeviceSharding or fully-replicated sharding is supported");15 }16 absl::Span<xla::ifrt::Device* const> ifrt_addressable_devices =17 sharding->devices()->AddressableDeviceList()->devices();1819 PjRtArray::PjRtBuffers buffers;20 buffers.reserve(ifrt_addressable_devices.size());21 for (xla::ifrt::Device* const device : ifrt_addressable_devices) {22 std::unique_ptr<PjRtBuffer> buffer;23 // ... memory-kind resolution trimmed ...24 ABSL_ASSIGN_OR_RETURN(xla::PjRtMemorySpace * memory_space,25 absl::down_cast<PjRtDevice*>(device)26 ->pjrt_device()27 ->default_memory_space());28 ABSL_ASSIGN_OR_RETURN(29 buffer,30 pjrt_client_->BufferFromHostBuffer(31 data, primitive_type, shape.dims(), byte_strides, semantics,32 on_done_with_host_buffer_per_device, memory_space, xla_layout));33 buffers.push_back(std::move(buffer));34 }35 return PjRtArray::Create(this, dtype, std::move(shape), std::move(sharding),36 std::move(buffers), std::move(pjrt_layout));37}38tsl::Future<LoadedExecutableRef> PjRtCompiler::CompileAndLoad(39 std::unique_ptr<Program> program, std::unique_ptr<CompileOptions> options) {40 if (!isa_and_nonnull<HloProgram>(program.get())) {41 return absl::InvalidArgumentError("PjRtCompiler requires an HloProgram");42 }43 std::unique_ptr<HloProgram> xla_program =44 cast<HloProgram>(std::move(program));45 ABSL_ASSIGN_OR_RETURN(auto xla_compile_options,46 GetXlaCompileOptions(std::move(options)));47 ABSL_RETURN_IF_ERROR(48 TranslateDeviceIds(client_, xla_compile_options->compile_options));49 return PjRtLoadedExecutable::Create(50 client_, std::move(*xla_program).ToMaybeOwningMlirModule(),51 std::move(xla_compile_options->compile_options) /* ... trimmed ... */);52}53absl::StatusOr<PjRtLoadedExecutable::ExecuteResult>54PjRtLoadedExecutable::Execute(absl::Span<ArrayRef> args,55 const ExecuteOptions& options,56 std::optional<DeviceListRef> devices) {57 std::vector<std::vector<PjRtBuffer*>> argument_handles;58 int num_computations = addressable_devices_.size();59 argument_handles.resize(num_computations);60 for (int i = 0; i < args.size(); ++i) {61 auto* pjrt_array = dyn_cast_or_null<PjRtCompatibleArray>(args[i].get());62 // ... shard-count check trimmed ...63 int j = 0;64 for (const auto& pjrt_buffer : pjrt_array->pjrt_buffers()) {65 argument_handles[j].push_back(pjrt_buffer.get());66 ++j;67 }68 }69 std::vector<std::vector<std::unique_ptr<PjRtBuffer>>> pjrt_outputs;70 std::optional<std::vector<tsl::Future<>>> returned_pjrt_futures;71 returned_pjrt_futures.emplace();72 ABSL_ASSIGN_OR_RETURN(pjrt_outputs,73 pjrt_loaded_executable_->Execute(argument_handles, opts,74 returned_pjrt_futures));75 status = JoinFutures(absl::MakeSpan(*returned_pjrt_futures));76 // ... then the reverse transpose, one output array per result ...77 outputs.push_back(*PjRtArray::Create(78 client_, output_dtypes_[i], output_shapes_[i], output_shardings_[i],79 std::move(buffers) /* ... */));80absl::StatusOr<xla::ifrt::ArrayRef> Array::MakeArrayFromHostBuffer(81 xla::ifrt::Client* client, std::shared_ptr<RpcHelper> rpc_helper,82 const void* data, DType dtype, Shape shape, /* ... */83 std::function<void()> on_done_with_host_buffer) {84 auto req = std::make_unique<MakeArrayFromHostBufferRequest>();85 dtype.ToProto(*req->mutable_dtype(), rpc_helper->ifrt_serdes_version());86 shape.ToProto(*req->mutable_shape(), rpc_helper->ifrt_serdes_version());87 ABSL_RETURN_IF_ERROR(sharding->ToProto(*req->mutable_sharding(),88 rpc_helper->ifrt_serdes_version()));89 // ... layout handling and host-buffer staging trimmed ...90 req->set_host_buffer_handle(host_buffer_handle);91 rpc_helper->MakeArrayFromHostBuffer(std::move(req));92 return xla::ifrt::ArrayRef(tsl::MakeRef<Array>(93 client, std::move(rpc_helper), dtype, std::move(shape),94 std::move(sharding), ArrayHandle{host_buffer_handle}, /* ... */));95}
01/06The IFRT promise, one array from host bytes with the sharding at construction, opens with two refusals. String dtypes detour to a dedicated path, and the sharding must be single-device or fully replicated: a genuinely sharded array never enters this function. The framework uploads each shard as its own single-device array and assembles them afterward, which is why this guard can afford to be strict.
1std::vector<ComputationClient::DataPtr> PjRtComputationClient::TransferToDevice(2 absl::Span<const std::shared_ptr<const TensorSource>> tensors) {3 // ... metrics and profiler scopes trimmed ...4 std::vector<ComputationClient::DataPtr> datas;5 datas.reserve(tensors.size());6 int64_t total_size = 0;7 for (auto& tensor : tensors) {8 xla::PjRtDevice* pjrt_device = StringToPjRtDevice(tensor->device());910 total_size += xla::ShapeUtil::ByteSizeOf(tensor->shape());1112 std::shared_ptr<xla::PjRtBuffer> buffer =13 std::move(client_14 ->BufferFromHostBuffer(15 tensor->data(), tensor->primitive_type(),16 tensor->dimensions(), tensor->byte_strides(),17 xla::PjRtClient::HostBufferSemantics::18 kImmutableUntilTransferCompletes,19 [tensor]() { /* frees tensor */ },20 *pjrt_device->default_memory_space(),21 /*device_layout=*/nullptr)22 .value());2324 ComputationClient::DataPtr data =25 std::make_shared<PjRtData>(tensor->device(), tensor->shape(), buffer);26 datas.push_back(data);27 }28 OutboundDataMetric()->AddSample(total_size);29 CreateDataHandlesCounter()->AddValue(datas.size());3031 return datas;32}33std::vector<ComputationClient::ComputationPtr> PjRtComputationClient::Compile(34 std::vector<ComputationClient::CompileInstance> instances) {35 // ... metrics and profiler scopes trimmed ...36 std::vector<ComputationClient::ComputationPtr> computations;37 // ... an env-flag static trimmed ...38 for (auto& instance : instances) {39 xla::CompileOptions compile_options;40 for (const auto& [name, value] : custom_compile_options_) {41 compile_options.env_option_overrides.push_back({name, value});42 }43 // ... the collective-matmul option overrides trimmed ...44 if (instance.is_sharded) {45 // TODO(yeounoh) multi-host, multi-slice configurations46 compile_options.executable_build_options.set_use_spmd_partitioning(true);47 // ... the output-propagation option trimmed ...48 int num_partitions = client_->device_count();49 compile_options.executable_build_options.set_num_partitions(50 num_partitions);51 compile_options.executable_build_options.set_num_replicas(1);52 // ... the tupling flag, then auto-SPMD options and logging trimmed ...53 // TODO(244391366) verify this is correct for the collectives ops54 xla::DeviceAssignment device_assignment(1, client_->device_count());55 // DeviceAssignment values must be the PjRtDevice ID, so we need to56 // unwind the global ordinal mapping.57 for (const auto& [device_id, global_ordinal] : global_ordinals_) {58 device_assignment(0, global_ordinal) = device_id;59 }60 compile_options.executable_build_options.set_device_assignment(61 device_assignment);62 } else {63 // ... an argument-layout TODO trimmed ...64 compile_options.executable_build_options.set_num_partitions(1);65 compile_options.executable_build_options.set_num_replicas(66 client_->device_count());67 // ... the same tupling flag trimmed ...68 xla::DeviceAssignment device_assignment(client_->device_count(), 1);69 // ... the same unwind comment trimmed ...70 for (const auto& [device_id, global_ordinal] : global_ordinals_) {71 device_assignment(global_ordinal, 0) = device_id;72 }73 compile_options.executable_build_options.set_device_assignment(74 device_assignment);75 }76 // ... a comment about raising a Python exception on compiler failure ...77 std::unique_ptr<xla::PjRtLoadedExecutable> executable;78 if (runtime::sys_util::GetEnvBool("XLA_STABLEHLO_COMPILE", false)) {79 // ... the HLO to StableHLO conversion and its CompileAndLoad trimmed ...80 } else {81 executable = util::RaisePythonValueErrorOnFailure([&] {82 return fake_xla_compile_ ? fake_xla_compile_()83 : client_->CompileAndLoad(instance.computation,84 compile_options);85 });86 }87 // ... compiled-memory-stats logging trimmed ...88 XLA_ASSIGN_OR_THROW(89 const std::vector<std::shared_ptr<xla::HloModule>>& hlo_modules,90 executable->GetHloModules());91 // ... an unused entry-computation local trimmed ...92 std::shared_ptr<PjRtComputation> pjrt_computation =93 std::make_shared<PjRtComputation>(94 std::move(xla::XlaComputation(hlo_modules[0]->ToProto())),95 instance.devices, std::move(executable));9697 computations.push_back(pjrt_computation);98 // ... a compile-handle counter trimmed ...99 }100101 return computations;102}103absl::StatusOr<std::vector<ComputationClient::DataPtr>>104PjRtComputationClient::ExecuteReplicated(105 const ComputationClient::Computation& computation,106 absl::Span<const ComputationClient::DataPtr> arguments,107 absl::Span<const std::string> devices,108 const ExecuteReplicatedOptions& options) {109 // Shared ownership of the timed section ensures that it will only get logged110 // once both `ExecuteReplicated` and the async work in `Execute` are111 // complete; a copy is held from the lambda that releases it when done.112 auto timed =113 std::make_shared<metrics::TimedSection>(ExecuteReplicatedMetric());114 // ... profiler scope trimmed ...115 const PjRtComputation& pjrt_computation =116 dynamic_cast<const PjRtComputation&>(computation);117118 std::vector<std::vector<xla::PjRtBuffer*>> argument_handles(119 devices.size(), std::vector<xla::PjRtBuffer*>(arguments.size()));120 {121 // ... a profiler scope, a counter, and the ParallelFor header trimmed ...122 for (int32_t i = start; i < end; ++i) {123 auto pjrt_data =124 std::dynamic_pointer_cast<PjRtShardedData>(arguments[i]);125 ABSL_CHECK_EQ(pjrt_data->shards.size(), devices.size())126 << "Expected one shard per device";127128 for (int32_t d = 0; d < devices.size(); d++) {129 std::shared_ptr<PjRtData> shard = pjrt_data->shards[d];130 // ... device address checks trimmed ...131 argument_handles[d][i] = shard->buffer.get();132 // ... the loops close and the counter joins ...133 }134135 xla::ExecuteOptions execute_options;136 execute_options.untuple_result = options.explode_tuple;137 execute_options.strict_shape_checking = true;138 // ... multi-slice and callback-layout options trimmed ...139 // Grab the shared lock and block the `WaitDeviceOps` until buffer is140 // ready. Since this is the SPMD code path. There is no points to grab141 // devices lock for every individual device.142 // ... a lock-timing log line trimmed ...143 auto op_tracker = operation_manager_.StartOperation(spmd_device_str);144 // ... its matching Done line trimmed ...145 std::optional<std::vector<xla::PjRtFuture<>>> returned_futures =146 std::vector<xla::PjRtFuture<>>();147 std::vector<std::vector<std::unique_ptr<xla::PjRtBuffer>>> results;148 {149 // ... a profiler scope trimmed ...150 XLA_ASSIGN_OR_RETURN(results, pjrt_computation.executable->Execute(151 std::move(argument_handles),152 execute_options, returned_futures));153154 (*returned_futures)[0].OnReady(155 std::move([timed, op_tracker = std::move(op_tracker)](156 absl::Status unused) mutable {157 timed.reset();158 TF_VLOG(3) << "ExecuteReplicated returned_future->OnReady finished";159 }));160 }161162 size_t num_outputs = results[0].size();163 std::vector<ComputationClient::DataPtr> data_handles(num_outputs);164 // ... output shapes and shardings resolved, then a ParallelFor over them ...165 std::vector<std::shared_ptr<PjRtData>> shards(devices.size());166 for (int32_t d = 0; d < devices.size(); d++) {167 std::unique_ptr<xla::PjRtBuffer> buffer =168 std::move(results[d][i]);169 shards[d] =170 std::make_shared<PjRtData>(devices[d], std::move(buffer));171 }172173 data_handles[i] = std::make_shared<PjRtShardedData>(174 spmd_device_str, output_shapes[i], std::move(shards),175 output_shardings[i]);176 // ... the pool joins ...177 return data_handles;178}179absl::StatusOr<std::vector<xla::Literal>>180PjRtComputationClient::TransferFromDevice(absl::Span<const DataPtr> handles) {181 // ... metrics and profiler scopes trimmed ...182 std::vector<xla::PjRtFuture<>> futures;183 futures.reserve(handles.size());184 std::vector<xla::Literal> literals;185 literals.reserve(handles.size());186 int64_t total_size = 0;187 for (auto handle : handles) {188 // Use XLA replication to reassemble the sharded data. If input handle189 // is not sharded, then it is a no-op.190 std::shared_ptr<PjRtData> pjrt_data = ReplicateShardedData(handle);191 ABSL_CHECK(pjrt_data) << "PjRt_data is null in " << __FUNCTION__;192 ABSL_CHECK(pjrt_data->buffer != nullptr)193 << "PjRt buffer is null in " << __FUNCTION__;194195 xla::Literal& literal = literals.emplace_back(196 xla::Literal(host_output_shape(pjrt_data->buffer.get())));197 futures.push_back(pjrt_data->buffer->ToLiteral(&literal));198199 total_size += literal.size_bytes();200 }201 XLA_RETURN_IF_ERROR(xla::JoinFutures(futures).Await());202 InboundDataMetric()->AddSample(total_size);203204 return literals;205}
01/09Parameters leave the host here, one tensor at a time. What crosses is a raw pointer plus the facts PJRT needs to read it: element type, dimensions, byte strides, and the memory space of the device that string named. The semantics argument is the promise about that pointer, kImmutableUntilTransferCompletes, so the runtime may read it until the copy finishes and the lambda then drops torch_xla's last reference to the tensor. The bytes were staged before this call, in AtenSource's constructor, which moved the at::Tensor to contiguous CPU memory in the target dtype. Underneath is the cpu-client walk's BufferFromHostBuffer, taking its copying branch rather than the zero-copy one, and a plugin catches the same call as PJRT_Client_BufferFromHostBuffer.
Readings
- pjrt_client.h ↗ the promises, primary source
- cpu_client.cc ↗ the worked example: CompileAndLoad down to RunHloPasses
- The PJRT C ABI ↗ the struct of function pointers, verbatim
- The PJRT-backed IFRT adapter ↗ way one, one file per wrapped class
- pjrt_computation_client.cpp ↗ the caller's side, one file, four functions
- The IFRT proxy ↗ way two, client and server halves