diff --git a/csrc/kernels/legacy/ibgda_device.cuh b/csrc/kernels/legacy/ibgda_device.cuh index 819bdd1c8..994fae67f 100644 --- a/csrc/kernels/legacy/ibgda_device.cuh +++ b/csrc/kernels/legacy/ibgda_device.cuh @@ -8,6 +8,8 @@ #pragma once +#include + #include #include #include @@ -78,11 +80,21 @@ __device__ static __forceinline__ nvshmemi_ibgda_device_state_t* ibgda_get_state return &nvshmemi_ibgda_device_state_d; } +template +__device__ static __forceinline__ nvshmemi_ibgda_device_qp_t* ibgda_get_rc_impl(StateType* state, int pe, int id) { + const auto num_rc_per_pe = state->num_rc_per_pe; + + if constexpr (std::is_same_v) { + return &state->globalmem + .rcs[pe * num_rc_per_pe * state->num_devices_initialized + id % (num_rc_per_pe * state->num_devices_initialized)]; + } else { + return &state->globalmem.rcs[pe + nvshmemi_device_state_d.npes * id]; + } +} + __device__ static __forceinline__ nvshmemi_ibgda_device_qp_t* ibgda_get_rc(int pe, int id) { auto state = ibgda_get_state(); - const auto num_rc_per_pe = ibgda_get_state()->num_rc_per_pe; - return &state->globalmem - .rcs[pe * num_rc_per_pe * state->num_devices_initialized + id % (num_rc_per_pe * state->num_devices_initialized)]; + return ibgda_get_rc_impl(state, pe, id); } __device__ static __forceinline__ void ibgda_lock_acquire(int* lock) { diff --git a/deep_ep/include/deep_ep/common/compiled.cuh b/deep_ep/include/deep_ep/common/compiled.cuh index 491e9307d..58d8bff9b 100644 --- a/deep_ep/include/deep_ep/common/compiled.cuh +++ b/deep_ep/include/deep_ep/common/compiled.cuh @@ -7,6 +7,13 @@ #define __CUDACC__ #endif +// Define __CUDACC_RDC__ so NVSHMEM device symbols use the correct extern declarations. +#ifndef DISABLE_NVSHMEM +#ifndef __CUDACC_RDC__ +#define __CUDACC_RDC__ // NOLINT(*-reserved-identifier) +#endif +#endif + // Remove Torch restrictions #ifdef __CUDA_NO_HALF_CONVERSIONS__ #undef __CUDA_NO_HALF_CONVERSIONS__