From 14d7ac37d342bb41eda45d967119f37f94d98c2c Mon Sep 17 00:00:00 2001 From: Hao Zhang Date: Tue, 1 Jul 2025 17:15:02 +0800 Subject: [PATCH] Add device guards in pytorch cuda kernels. --- qmb/_hamiltonian_cuda.cu | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/qmb/_hamiltonian_cuda.cu b/qmb/_hamiltonian_cuda.cu index 2c208e3..c561e7e 100644 --- a/qmb/_hamiltonian_cuda.cu +++ b/qmb/_hamiltonian_cuda.cu @@ -1,4 +1,5 @@ #include +#include #include #include #include @@ -180,6 +181,7 @@ auto apply_within_interface( std::int64_t batch_size = configs.size(0); std::int64_t result_batch_size = result_configs.size(0); std::int64_t term_number = site.size(0); + at::cuda::CUDAGuard cuda_device_guard(device_id); TORCH_CHECK(configs.device().type() == torch::kCUDA, "configs must be on CUDA.") TORCH_CHECK(configs.device().index() == device_id, "configs must be on the same device as others."); @@ -560,6 +562,7 @@ auto find_relative_interface( std::int64_t batch_size = configs.size(0); std::int64_t term_number = site.size(0); std::int64_t exclude_size = exclude_configs.size(0); + at::cuda::CUDAGuard cuda_device_guard(device_id); TORCH_CHECK(configs.device().type() == torch::kCUDA, "configs must be on CUDA.") TORCH_CHECK(configs.device().index() == device_id, "configs must be on the same device as others."); @@ -779,6 +782,7 @@ auto single_relative_interface(const torch::Tensor& configs, const torch::Tensor std::int64_t device_id = configs.device().index(); std::int64_t batch_size = configs.size(0); std::int64_t term_number = site.size(0); + at::cuda::CUDAGuard cuda_device_guard(device_id); TORCH_CHECK(configs.device().type() == torch::kCUDA, "configs must be on CUDA.") TORCH_CHECK(configs.device().index() == device_id, "configs must be on the same device as others.");