From 2016b1f5ce2eb6a538fa119e13ed23d644a7d170 Mon Sep 17 00:00:00 2001 From: Hao Zhang Date: Mon, 16 Jun 2025 23:34:01 +0800 Subject: [PATCH] Use `cuda::std::array` instead of `std::array` in cuda kernels. --- qmb/_hamiltonian_cuda.cu | 171 +++++++++++++++++++-------------------- 1 file changed, 85 insertions(+), 86 deletions(-) diff --git a/qmb/_hamiltonian_cuda.cu b/qmb/_hamiltonian_cuda.cu index c54647d..b8099ec 100644 --- a/qmb/_hamiltonian_cuda.cu +++ b/qmb/_hamiltonian_cuda.cu @@ -11,7 +11,7 @@ constexpr torch::DeviceType device = torch::kCUDA; template struct array_less { - __device__ bool operator()(const std::array& lhs, const std::array& rhs) const { + __device__ bool operator()(const cuda::std::array& lhs, const cuda::std::array& rhs) const { for (std::int64_t i = 0; i < size; ++i) { if (lhs[i] < rhs[i]) { return true; @@ -26,14 +26,14 @@ struct array_less { template struct array_square_greater { - __device__ T square(const std::array& value) const { + __device__ T square(const cuda::std::array& value) const { T result = 0; for (std::int64_t i = 0; i < size; ++i) { result += value[i] * value[i]; } return result; } - __device__ bool operator()(const std::array& lhs, const std::array& rhs) const { + __device__ bool operator()(const cuda::std::array& lhs, const cuda::std::array& rhs) const { return square(lhs) > square(rhs); } }; @@ -52,11 +52,11 @@ __device__ void set_bit(std::uint8_t* data, std::uint8_t index, bool value) { template __device__ std::pair hamiltonian_apply_kernel( - std::array& current_configs, + cuda::std::array& current_configs, std::int64_t term_index, std::int64_t batch_index, - const std::array* site, // term_number - const std::array* kind // term_number + const cuda::std::array* site, // term_number + const cuda::std::array* kind // term_number ) { static_assert(particle_cut == 1 || particle_cut == 2, "particle_cut != 1 or 2 not implemented"); bool success = true; @@ -89,15 +89,15 @@ __device__ void apply_within_kernel( std::int64_t term_number, std::int64_t batch_size, std::int64_t result_batch_size, - const std::array* site, // term_number - const std::array* kind, // term_number - const std::array* coef, // term_number - const std::array* configs, // batch_size - const std::array* psi, // batch_size - const std::array* result_configs, // result_batch_size - std::array* result_psi + const cuda::std::array* site, // term_number + const cuda::std::array* kind, // term_number + const cuda::std::array* coef, // term_number + const cuda::std::array* configs, // batch_size + const cuda::std::array* psi, // batch_size + const cuda::std::array* result_configs, // result_batch_size + cuda::std::array* result_psi ) { - std::array current_configs = configs[batch_index]; + cuda::std::array current_configs = configs[batch_index]; auto [success, parity] = hamiltonian_apply_kernel( /*current_configs=*/current_configs, /*term_index=*/term_index, @@ -138,13 +138,13 @@ __global__ void apply_within_kernel_interface( std::int64_t term_number, std::int64_t batch_size, std::int64_t result_batch_size, - const std::array* site, // term_number - const std::array* kind, // term_number - const std::array* coef, // term_number - const std::array* configs, // batch_size - const std::array* psi, // batch_size - const std::array* result_configs, // result_batch_size - std::array* result_psi + const cuda::std::array* site, // term_number + const cuda::std::array* kind, // term_number + const cuda::std::array* coef, // term_number + const cuda::std::array* configs, // batch_size + const cuda::std::array* psi, // batch_size + const cuda::std::array* result_configs, // result_batch_size + cuda::std::array* result_psi ) { std::int64_t term_index = blockIdx.x * blockDim.x + threadIdx.x; std::int64_t batch_index = blockIdx.y * blockDim.y + threadIdx.y; @@ -242,8 +242,8 @@ auto apply_within_interface( thrust::sort_by_key( policy, - reinterpret_cast*>(sorted_result_configs.data_ptr()), - reinterpret_cast*>(sorted_result_configs.data_ptr()) + result_batch_size, + reinterpret_cast*>(sorted_result_configs.data_ptr()), + reinterpret_cast*>(sorted_result_configs.data_ptr()) + result_batch_size, reinterpret_cast(result_sort_index.data_ptr()), array_less() ); @@ -256,13 +256,13 @@ auto apply_within_interface( /*term_number=*/term_number, /*batch_size=*/batch_size, /*result_batch_size=*/result_batch_size, - /*site=*/reinterpret_cast*>(site.data_ptr()), - /*kind=*/reinterpret_cast*>(kind.data_ptr()), - /*coef=*/reinterpret_cast*>(coef.data_ptr()), - /*configs=*/reinterpret_cast*>(configs.data_ptr()), - /*psi=*/reinterpret_cast*>(psi.data_ptr()), - /*result_configs=*/reinterpret_cast*>(sorted_result_configs.data_ptr()), - /*result_psi=*/reinterpret_cast*>(sorted_result_psi.data_ptr()) + /*site=*/reinterpret_cast*>(site.data_ptr()), + /*kind=*/reinterpret_cast*>(kind.data_ptr()), + /*coef=*/reinterpret_cast*>(coef.data_ptr()), + /*configs=*/reinterpret_cast*>(configs.data_ptr()), + /*psi=*/reinterpret_cast*>(psi.data_ptr()), + /*result_configs=*/reinterpret_cast*>(sorted_result_configs.data_ptr()), + /*result_psi=*/reinterpret_cast*>(sorted_result_psi.data_ptr()) ); AT_CUDA_CHECK(cudaStreamSynchronize(stream)); @@ -426,7 +426,7 @@ __device__ void add_into_heap(T* heap, int* mutex, std::int64_t heap_size, const template struct array_first_double_less { - __device__ double first_double(const std::array& value) const { + __device__ double first_double(const cuda::std::array& value) const { double result; for (std::int64_t i = 0; i < sizeof(double); ++i) { reinterpret_cast(&result)[i] = reinterpret_cast(&value[0])[i]; @@ -434,8 +434,10 @@ struct array_first_double_less { return result; } - __device__ bool - operator()(const std::array& lhs, const std::array& rhs) const { + __device__ bool operator()( + const cuda::std::array& lhs, + const cuda::std::array& rhs + ) const { return first_double(lhs) < first_double(rhs); } }; @@ -447,17 +449,17 @@ __device__ void find_relative_kernel( std::int64_t term_number, std::int64_t batch_size, std::int64_t exclude_size, - const std::array* site, // term_number - const std::array* kind, // term_number - const std::array* coef, // term_number - const std::array* configs, // batch_size - const std::array* psi, // batch_size - const std::array* exclude_configs, // exclude_size - std::array* heap, + const cuda::std::array* site, // term_number + const cuda::std::array* kind, // term_number + const cuda::std::array* coef, // term_number + const cuda::std::array* configs, // batch_size + const cuda::std::array* psi, // batch_size + const cuda::std::array* exclude_configs, // exclude_size + cuda::std::array* heap, int* mutex, std::int64_t heap_size ) { - std::array current_configs = configs[batch_index]; + cuda::std::array current_configs = configs[batch_index]; auto [success, parity] = hamiltonian_apply_kernel( /*current_configs=*/current_configs, /*term_index=*/term_index, @@ -493,19 +495,16 @@ __device__ void find_relative_kernel( double imag = sign * (coef[term_index][0] * psi[batch_index][1] + coef[term_index][1] * psi[batch_index][0]); // Currently, the weight is calculated as the probability of the state, but it can be changed to other values in the future. double weight = real * real + imag * imag; - std::array value; + cuda::std::array value; for (std::int64_t i = 0; i < sizeof(double) / sizeof(uint8_t); ++i) { value[i] = reinterpret_cast(&weight)[i]; } for (std::int64_t i = 0; i < n_qubytes; ++i) { value[i + sizeof(double) / sizeof(uint8_t)] = current_configs[i]; } - add_into_heap, array_first_double_less>( - heap, - mutex, - heap_size, - value - ); + add_into_heap< + cuda::std::array, + array_first_double_less>(heap, mutex, heap_size, value); } template @@ -513,13 +512,13 @@ __global__ void find_relative_kernel_interface( std::int64_t term_number, std::int64_t batch_size, std::int64_t exclude_size, - const std::array* site, // term_number - const std::array* kind, // term_number - const std::array* coef, // term_number - const std::array* configs, // batch_size - const std::array* psi, // batch_size - const std::array* exclude_configs, // exclude_size - std::array* heap, + const cuda::std::array* site, // term_number + const cuda::std::array* kind, // term_number + const cuda::std::array* coef, // term_number + const cuda::std::array* configs, // batch_size + const cuda::std::array* psi, // batch_size + const cuda::std::array* exclude_configs, // exclude_size + cuda::std::array* heap, int* mutex, std::int64_t heap_size ) { @@ -620,8 +619,8 @@ auto find_relative_interface( thrust::sort( policy, - reinterpret_cast*>(sorted_exclude_configs.data_ptr()), - reinterpret_cast*>(sorted_exclude_configs.data_ptr()) + exclude_size, + reinterpret_cast*>(sorted_exclude_configs.data_ptr()), + reinterpret_cast*>(sorted_exclude_configs.data_ptr()) + exclude_size, array_less() ); @@ -629,8 +628,8 @@ auto find_relative_interface( {count_selected, n_qubytes + sizeof(double) / sizeof(std::uint8_t)}, torch::TensorOptions().dtype(torch::kUInt8).device(device, device_id) ); - std::array* heap = - reinterpret_cast*>(result_pool.data_ptr()); + cuda::std::array* heap = + reinterpret_cast*>(result_pool.data_ptr()); int* mutex; AT_CUDA_CHECK(cudaMalloc(&mutex, sizeof(int) * count_selected)); AT_CUDA_CHECK(cudaMemset(mutex, 0, sizeof(int) * count_selected)); @@ -643,12 +642,12 @@ auto find_relative_interface( /*term_number=*/term_number, /*batch_size=*/batch_size, /*exclude_size=*/exclude_size, - /*site=*/reinterpret_cast*>(site.data_ptr()), - /*kind=*/reinterpret_cast*>(kind.data_ptr()), - /*coef=*/reinterpret_cast*>(coef.data_ptr()), - /*configs=*/reinterpret_cast*>(configs.data_ptr()), - /*psi=*/reinterpret_cast*>(psi.data_ptr()), - /*exclude_configs=*/reinterpret_cast*>(sorted_exclude_configs.data_ptr()), + /*site=*/reinterpret_cast*>(site.data_ptr()), + /*kind=*/reinterpret_cast*>(kind.data_ptr()), + /*coef=*/reinterpret_cast*>(coef.data_ptr()), + /*configs=*/reinterpret_cast*>(configs.data_ptr()), + /*psi=*/reinterpret_cast*>(psi.data_ptr()), + /*exclude_configs=*/reinterpret_cast*>(sorted_exclude_configs.data_ptr()), /*heap=*/heap, /*mutex=*/mutex, /*heap_size=*/count_selected @@ -679,16 +678,16 @@ __device__ void single_relative_kernel( std::int64_t batch_size, std::int64_t exclude_size, std::uint64_t seed, - const std::array* site, // term_number - const std::array* kind, // term_number - const std::array* coef, // term_number - const std::array* configs, // batch_size - const std::array* exclude_configs, // exclude_size - std::array* result_configs, // batch_size + const cuda::std::array* site, // term_number + const cuda::std::array* kind, // term_number + const cuda::std::array* coef, // term_number + const cuda::std::array* configs, // batch_size + const cuda::std::array* exclude_configs, // exclude_size + cuda::std::array* result_configs, // batch_size double* score, // batch_size int* mutex // batch_size ) { - std::array current_configs = configs[batch_index]; + cuda::std::array current_configs = configs[batch_index]; auto [success, parity] = hamiltonian_apply_kernel( /*current_configs=*/current_configs, /*term_index=*/term_index, @@ -741,12 +740,12 @@ __global__ void single_relative_kernel_interface( std::int64_t batch_size, std::int64_t exclude_size, std::uint64_t seed, - const std::array* site, // term_number - const std::array* kind, // term_number - const std::array* coef, // term_number - const std::array* configs, // batch_size - const std::array* exclude_configs, // exclude_size - std::array* result_configs, // batch_size + const cuda::std::array* site, // term_number + const cuda::std::array* kind, // term_number + const cuda::std::array* coef, // term_number + const cuda::std::array* configs, // batch_size + const cuda::std::array* exclude_configs, // exclude_size + cuda::std::array* result_configs, // batch_size double* score, // batch_size int* mutex // batch_size ) { @@ -823,8 +822,8 @@ auto single_relative_interface(const torch::Tensor& configs, const torch::Tensor thrust::sort( policy, - reinterpret_cast*>(sorted_configs.data_ptr()), - reinterpret_cast*>(sorted_configs.data_ptr()) + batch_size, + reinterpret_cast*>(sorted_configs.data_ptr()), + reinterpret_cast*>(sorted_configs.data_ptr()) + batch_size, array_less() ); @@ -847,12 +846,12 @@ auto single_relative_interface(const torch::Tensor& configs, const torch::Tensor /*batch_size=*/batch_size, /*exclude_size=*/batch_size, /*seed=*/seed, - /*site=*/reinterpret_cast*>(site.data_ptr()), - /*kind=*/reinterpret_cast*>(kind.data_ptr()), - /*coef=*/reinterpret_cast*>(coef.data_ptr()), - /*configs=*/reinterpret_cast*>(configs.data_ptr()), - /*exclude_configs=*/reinterpret_cast*>(sorted_configs.data_ptr()), - /*result_configs=*/reinterpret_cast*>(result_configs.data_ptr()), + /*site=*/reinterpret_cast*>(site.data_ptr()), + /*kind=*/reinterpret_cast*>(kind.data_ptr()), + /*coef=*/reinterpret_cast*>(coef.data_ptr()), + /*configs=*/reinterpret_cast*>(configs.data_ptr()), + /*exclude_configs=*/reinterpret_cast*>(sorted_configs.data_ptr()), + /*result_configs=*/reinterpret_cast*>(result_configs.data_ptr()), /*score=*/reinterpret_cast(score.data_ptr()), /*mutex=*/mutex );