Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 15 additions & 32 deletions cpp/src/torch_tensorrt/executorch/TensorRTBackend.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
#include <cuda_runtime.h>

#include <executorch/extension/cuda/caller_stream.h>
#include <executorch/extension/cuda/device_guard.h>
#include <executorch/runtime/backend/interface.h>
#include <executorch/runtime/core/exec_aten/util/tensor_util.h>
#include <executorch/runtime/platform/log.h>
Expand Down Expand Up @@ -78,7 +79,11 @@ void TRTLogger::log(Severity severity, const char* msg) noexcept {
}

EngineHandle::~EngineHandle() {
cudaSetDevice(device_id);
// Freeing runs on the engine's device. Borrowed rather than selected outright,
// because this can run from arena teardown on a thread that was working
// elsewhere, and a destructor has no way to report a failure.
const auto guard = ::executorch::extension::cuda::CUDAGuard::create(device_id);
(void)guard;
// No wait here: execute already waited, and a device-wide one blocks unrelated work.
for (void* p : cached_input_ptrs) {
if (p != nullptr) {
Expand Down Expand Up @@ -328,16 +333,15 @@ Result<DelegateHandle*> TensorRTBackend::init(
handle->output_binding_names = std::move(header.output_binding_names);
handle->device_id = header.device_id;

cudaError_t cuda_err = cudaSetDevice(handle->device_id);
if (cuda_err != cudaSuccess) {
ET_LOG(
Error, "TensorRTBackend::init: cudaSetDevice(%d) failed: %s", handle->device_id, cudaGetErrorString(cuda_err));
auto device_guard = ::executorch::extension::cuda::CUDAGuard::create(handle->device_id);
if (!device_guard.ok()) {
ET_LOG(Error, "TensorRTBackend::init: cannot select device %d", handle->device_id);
return Error::InvalidProgram;
}

// Whether this device can reach pageable host memory at all. Speed is the query below.
int pageable_access = 0;
cuda_err = cudaDeviceGetAttribute(&pageable_access, cudaDevAttrPageableMemoryAccess, handle->device_id);
cudaError_t cuda_err = cudaDeviceGetAttribute(&pageable_access, cudaDevAttrPageableMemoryAccess, handle->device_id);
if (cuda_err != cudaSuccess) {
ET_LOG(
Info,
Expand Down Expand Up @@ -649,36 +653,15 @@ Error TensorRTBackend::execute(BackendExecutionContext& context, DelegateHandle*
return Error::InvalidArgument;
}

int entry_device = -1;
cudaError_t cuda_err = cudaGetDevice(&entry_device);
if (cuda_err != cudaSuccess) {
ET_LOG(Error, "TensorRTBackend::execute: cudaGetDevice failed: %s", cudaGetErrorString(cuda_err));
return Error::InvalidProgram;
}
// Put the engine on its own device for multi-GPU correctness, restoring the
// caller's device on exit; green-context confinement rides the selected stream,
// independent of the current device/context.
const bool switch_device = (entry_device != engine->device_id);
if (switch_device) {
cuda_err = cudaSetDevice(engine->device_id);
if (cuda_err != cudaSuccess) {
ET_LOG(
Error,
"TensorRTBackend::execute: cudaSetDevice(%d) failed: %s",
engine->device_id,
cudaGetErrorString(cuda_err));
return Error::InvalidProgram;
}
auto device_guard = ::executorch::extension::cuda::CUDAGuard::create(engine->device_id);
if (!device_guard.ok()) {
ET_LOG(Error, "TensorRTBackend::execute: cannot select device %d", engine->device_id);
return Error::InvalidProgram;
}
struct DeviceRestore {
int device;
bool active;
~DeviceRestore() {
if (active) {
cudaSetDevice(device);
}
}
} device_restore{entry_device, switch_device};
cudaError_t cuda_err = cudaSuccess;

std::unique_lock<std::mutex> lock(engine->mu);

Expand Down
71 changes: 71 additions & 0 deletions tests/py/dynamo/executorch/test_cuda_partitioner_composition.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: BSD-3-Clause

import gc
import importlib.util
import os
import shutil
Expand Down Expand Up @@ -390,3 +391,73 @@ def test_the_partitioner_refuses_a_device_it_cannot_run_on(requested, accepted):
else:
with pytest.raises(ValueError, match="not a device this delegate runs on"):
TensorRTPartitioner(compile_specs=specs)


def test_the_caller_keeps_its_device_across_load_run_and_free(tmp_path):
"""A program whose engine sits on another card must not move the caller.

Loading, running and freeing all select the engine's device. Each one has to put
the caller's device back, because a coalesced program holds several delegates on
several cards and the thread that loads one is not doing that delegate's work.
Needs two cards: with one there is nothing to switch to, so nothing to restore.
"""
if torch.cuda.device_count() < 2:
pytest.skip(
"needs two CUDA devices to tell a restored device from an unchanged one"
)

import torch_tensorrt

class Model(torch.nn.Module):
def forward(self, x):
return torch.cos(torch.erfinv(torch.tanh(x)))

engine_device = 1
caller_device = 0

with torch.cuda.device(engine_device):
model = Model().eval().to(f"cuda:{engine_device}")
inputs = (torch.randn(64, 64, device=f"cuda:{engine_device}"),)
exported = torch.export.export(model, inputs)
trt_gm = torch_tensorrt.dynamo.compile(
exported, inputs=list(inputs), min_block_size=1, truncate_double=True
)
out = tmp_path / "other_card.pte"
torch_tensorrt.save(
trt_gm,
str(out),
output_format="executorch",
retrace=False,
arg_inputs=list(inputs),
partitioners=[_cuda_partitioner()],
)

delegate_ids = _delegate_ids(out)
assert (
"TensorRTBackend" in delegate_ids
), f"nothing went to TensorRT; {delegate_ids}"

import torch_tensorrt_executorch_runtime # noqa: F401
from executorch.runtime import Runtime

torch.cuda.set_device(caller_device)

program = Runtime.get().load_program(out)
assert (
torch.cuda.current_device() == caller_device
), "loading moved the caller's device"

method = program.load_method("forward")
assert torch.cuda.current_device() == caller_device, "preparing the method moved it"

method.execute((torch.randn(64, 64, device=f"cuda:{engine_device}"),))
assert (
torch.cuda.current_device() == caller_device
), "running moved the caller's device"

del method
del program
gc.collect()
assert (
torch.cuda.current_device() == caller_device
), "freeing moved the caller's device"
26 changes: 19 additions & 7 deletions third_party/executorch/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -60,19 +60,23 @@ cmake(
# build the implementation once as a cc_binary(linkshared=True), re-import the
# .so via cc_import, and expose it through a srcs-less cc_library that also
# carries the headers (with include remap) and the cudart dep. Consumers depend
# on ":extension_cuda"; caller_stream.cpp is only ever linked into the one .so.
# on ":extension_cuda"; the sources are only ever linked into the one .so.

# Private implementation. It carries hdrs + include_prefix/strip_include_prefix
# so caller_stream.cpp can find its own <executorch/extension/cuda/...> header
# so each source can find its own <executorch/extension/cuda/...> header
# (the physical tree has an extra executorch/ dir; a cc_binary's srcs headers
# get no such remap). alwayslink=True guarantees the object -- and thus
# getCallerStream/CallerStreamGuard -- is pulled into the shared object below,
# even though that source-less cc_binary references none of its symbols.
# get no such remap). alwayslink=True guarantees the objects -- and thus the
# caller-stream and device-guard symbols -- are pulled into the shared object
# below, even though that source-less cc_binary references none of them.
cc_library(
name = "extension_cuda_impl",
srcs = ["executorch/extension/cuda/caller_stream.cpp"],
srcs = [
"executorch/extension/cuda/caller_stream.cpp",
"executorch/extension/cuda/device_guard.cpp",
],
hdrs = [
"executorch/extension/cuda/caller_stream.h",
"executorch/extension/cuda/device_guard.h",
"executorch/extension/cuda/export.h",
],
local_defines = [
Expand All @@ -84,7 +88,15 @@ cc_library(
alwayslink = True,
linkstatic = True,
visibility = ["//visibility:private"],
deps = ["@cuda//:cudart"],
# device_guard.cpp returns an Error and logs, so it needs the runtime's headers,
# which caller_stream.cpp beside it does not. Only the headers: the logging symbols
# themselves resolve at load time from the runtime the consumer already loads.
# Linking the static core here instead would give the process a second copy of
# every registry it carries.
deps = [
":executorch_headers",
"@cuda//:cudart",
],
)

# The single shared object. The explicit -soname pins DT_SONAME to the bare
Expand Down
Loading