From 7015997a09288e3aca3f4d66d682665deb3b1fda Mon Sep 17 00:00:00 2001 From: qinyiqun Date: Mon, 3 Aug 2026 06:04:54 +0000 Subject: [PATCH 1/2] fix(nvidia): support graph replay, MoE decode, and diff builds --- .../elementwise/nvidia/elementwise_nvidia.cuh | 39 ++++++++++++++++++- src/infiniop/ops/diff/nvidia/diff_nvidia.cuh | 3 ++ .../nvidia/moe_fused_dense_nvidia.cu | 19 ++++++--- 3 files changed, 54 insertions(+), 7 deletions(-) diff --git a/src/infiniop/elementwise/nvidia/elementwise_nvidia.cuh b/src/infiniop/elementwise/nvidia/elementwise_nvidia.cuh index f8cb457ec..ae2d62422 100644 --- a/src/infiniop/elementwise/nvidia/elementwise_nvidia.cuh +++ b/src/infiniop/elementwise/nvidia/elementwise_nvidia.cuh @@ -80,6 +80,28 @@ __device__ __forceinline__ const T *typedInputPtr(const void *ptr) { return reinterpret_cast(ptr); } +template +struct InputPointerArray { + const void *values[N]; +}; + +/** + * @brief Stores elementwise input pointers in device workspace. + * + * The pointer array is passed by value as a kernel argument. This is required + * for CUDA Graph capture: a captured cudaMemcpyAsync from inputs.data() would + * retain a pointer to a temporary host std::vector that is destroyed before + * graph replay. + */ +template +INFINIOP_CUDA_KERNEL storeInputPointers( + const void **output, + InputPointerArray inputs) { + for (size_t i = threadIdx.x; i < N; i += blockDim.x) { + output[i] = inputs.values[i]; + } +} + /** * @brief Computes the output index in memory, accounting for strides if non-contiguous. * @@ -580,8 +602,21 @@ private: const int8_t *info_meta_start = info.getMetaStart(); const int8_t *d_meta_start = reinterpret_cast(workspace) + input_arr_size; - // copy the input pointer array and meta to device - CHECK_CUDA(cudaMemcpyAsync(workspace, h_inputs_arr, input_arr_size, cudaMemcpyHostToDevice, stream)); + cudaStreamCaptureStatus capture_status = cudaStreamCaptureStatusNone; + CHECK_CUDA(cudaStreamIsCapturing(stream, &capture_status)); + if (capture_status == cudaStreamCaptureStatusNone) { + CHECK_CUDA(cudaMemcpyAsync(workspace, h_inputs_arr, input_arr_size, cudaMemcpyHostToDevice, stream)); + } else { + // A captured H2D copy from a temporary std::vector would retain an + // invalid host pointer for replay. Kernel arguments are stored by + // value in the graph node instead. + InputPointerArray input_pointers{}; + for (size_t i = 0; i < N; ++i) { + input_pointers.values[i] = h_inputs_arr[i]; + } + storeInputPointers<<<1, N, 0, stream>>>( + reinterpret_cast(workspace), input_pointers); + } CHECK_CUDA(cudaMemcpyAsync((void *)d_meta_start, info_meta_start, info.getMetaMemSize(), cudaMemcpyHostToDevice, stream)); // offset/assign the pointers diff --git a/src/infiniop/ops/diff/nvidia/diff_nvidia.cuh b/src/infiniop/ops/diff/nvidia/diff_nvidia.cuh index 83772d853..24f5a958d 100644 --- a/src/infiniop/ops/diff/nvidia/diff_nvidia.cuh +++ b/src/infiniop/ops/diff/nvidia/diff_nvidia.cuh @@ -1,8 +1,11 @@ #ifndef __DIFF_NVIDIA_H__ #define __DIFF_NVIDIA_H__ +#include "../../../../utils.h" #include "../../../operator.h" #include +#include +#include namespace op::diff::nvidia { diff --git a/src/infiniop/ops/moe_fused_dense/nvidia/moe_fused_dense_nvidia.cu b/src/infiniop/ops/moe_fused_dense/nvidia/moe_fused_dense_nvidia.cu index 0b6357f9f..2df6a74f0 100644 --- a/src/infiniop/ops/moe_fused_dense/nvidia/moe_fused_dense_nvidia.cu +++ b/src/infiniop/ops/moe_fused_dense/nvidia/moe_fused_dense_nvidia.cu @@ -403,6 +403,7 @@ infiniStatus_t launch_cutlass_gemm_grouped_device_meta(int problem_count, int64_t *d_ldb, int64_t *d_ldc, int64_t *d_ldd, + bool full_occupancy_wave, cudaStream_t stream) { if (problem_count == 0) { return INFINI_STATUS_SUCCESS; @@ -436,7 +437,15 @@ infiniStatus_t launch_cutlass_gemm_grouped_device_meta(int problem_count, cutlass::gemm::kernel::GroupScheduleMode::kDeviceOnly>::GemmKernel; using Gemm = cutlass::gemm::device::GemmGrouped; - const int threadblock_count = std::min(Gemm::sufficient(), problem_count); + const int occupancy_wave = Gemm::sufficient(); + // Decode routes to only top-k experts, but each expert still contains many + // N-dimension tiles. A full persistent wave lets CUTLASS distribute those + // tiles across the GPU instead of leaving most SMs idle. Prefill retains + // the previous problem-count cap to avoid scheduler overhead from excess + // threadblocks when most experts receive no tokens. + const int threadblock_count = full_occupancy_wave + ? occupancy_wave + : std::min(occupancy_wave, problem_count); if (threadblock_count <= 0) { return INFINI_STATUS_DEVICE_ARCHITECTURE_NOT_SUPPORTED; } @@ -547,7 +556,7 @@ infiniStatus_t calculate_typed(const MoeFusedDenseInfo &info, // A D2H count copy and stream sync cannot be captured for replay. CHECK_STATUS(launch_cutlass_gemm_grouped_device_meta( topk, grouped_problems, grouped_ptr_a, grouped_ptr_b, grouped_ptr_c, grouped_ptr_d, - grouped_lda, grouped_ldb, grouped_ldc, grouped_ldd, stream)); + grouped_lda, grouped_ldb, grouped_ldc, grouped_ldd, true, stream)); swiglu_kernel<<<(topk * intermediate_size + 255) / 256, 256, 0, stream>>>(gate_up, activated, topk, intermediate_size); @@ -560,7 +569,7 @@ infiniStatus_t calculate_typed(const MoeFusedDenseInfo &info, output_permutation, pairs, num_experts, hidden_size, intermediate_size, block_size); CHECK_STATUS(launch_cutlass_gemm_grouped_device_meta( topk, grouped_problems, grouped_ptr_a, grouped_ptr_b, grouped_ptr_c, grouped_ptr_d, - grouped_lda, grouped_ldb, grouped_ldc, grouped_ldd, stream)); + grouped_lda, grouped_ldb, grouped_ldc, grouped_ldd, true, stream)); apply_shuffle_mul_sum_kernel<<<1, std::min(hidden_size, 1024), 0, stream>>>( expert_out, reinterpret_cast(output), output_permutation, @@ -604,7 +613,7 @@ infiniStatus_t calculate_typed(const MoeFusedDenseInfo &info, // avoids making the runtime problem count depend on device routing data. CHECK_STATUS(launch_cutlass_gemm_grouped_device_meta( num_experts, grouped_problems, grouped_ptr_a, grouped_ptr_b, grouped_ptr_c, grouped_ptr_d, - grouped_lda, grouped_ldb, grouped_ldc, grouped_ldd, stream)); + grouped_lda, grouped_ldb, grouped_ldc, grouped_ldd, false, stream)); swiglu_kernel<<<(max_num_tokens_padded * intermediate_size + 255) / 256, 256, 0, stream>>>( gate_up, activated, max_num_tokens_padded, intermediate_size); @@ -615,7 +624,7 @@ infiniStatus_t calculate_typed(const MoeFusedDenseInfo &info, w2_t, expert_out, num_experts, hidden_size, intermediate_size); CHECK_STATUS(launch_cutlass_gemm_grouped_device_meta( num_experts, grouped_problems, grouped_ptr_a, grouped_ptr_b, grouped_ptr_c, grouped_ptr_d, - grouped_lda, grouped_ldb, grouped_ldc, grouped_ldd, stream)); + grouped_lda, grouped_ldb, grouped_ldc, grouped_ldd, false, stream)); apply_shuffle_mul_sum_kernel<<>>( expert_out, reinterpret_cast(output), output_permutation, From a4e548d95bf530c7e4de633a0415dc3aaba0c9ca Mon Sep 17 00:00:00 2001 From: wooway777 Date: Fri, 14 Aug 2026 22:29:28 +0800 Subject: [PATCH 2/2] fix(nvidia): resolve elementwise build after rebase --- .../elementwise/nvidia/elementwise_nvidia.cuh | 174 ++++++------------ .../nvidia/elementwise_nvidia_api.cuh | 4 +- .../ops/hardtanh/nvidia/hardtanh_nvidia.cu | 4 +- .../mul_scalar/nvidia/mul_scalar_nvidia.cu | 8 +- 4 files changed, 67 insertions(+), 123 deletions(-) diff --git a/src/infiniop/elementwise/nvidia/elementwise_nvidia.cuh b/src/infiniop/elementwise/nvidia/elementwise_nvidia.cuh index ae2d62422..d93abf06e 100644 --- a/src/infiniop/elementwise/nvidia/elementwise_nvidia.cuh +++ b/src/infiniop/elementwise/nvidia/elementwise_nvidia.cuh @@ -80,28 +80,6 @@ __device__ __forceinline__ const T *typedInputPtr(const void *ptr) { return reinterpret_cast(ptr); } -template -struct InputPointerArray { - const void *values[N]; -}; - -/** - * @brief Stores elementwise input pointers in device workspace. - * - * The pointer array is passed by value as a kernel argument. This is required - * for CUDA Graph capture: a captured cudaMemcpyAsync from inputs.data() would - * retain a pointer to a temporary host std::vector that is destroyed before - * graph replay. - */ -template -INFINIOP_CUDA_KERNEL storeInputPointers( - const void **output, - InputPointerArray inputs) { - for (size_t i = threadIdx.x; i < N; i += blockDim.x) { - output[i] = inputs.values[i]; - } -} - /** * @brief Computes the output index in memory, accounting for strides if non-contiguous. * @@ -194,23 +172,22 @@ INFINIOP_CUDA_KERNEL elementwiseKernel( const ptrdiff_t *__restrict__ output_strides, const ptrdiff_t *__restrict__ input_strides, Tdata *output, - const void *const *inputs, + InputPointerArray inputs, size_t offset, Args... args) { size_t idx = blockIdx.x * blockDim.x + threadIdx.x + offset; if (idx < output_size) { - const Tdata *const *typed_inputs = reinterpret_cast(inputs); size_t out_idx = getOutputIndex(idx, output_contiguous, ndim, output_shape, output_strides); InputIndexer indexer{idx, ndim, input_contiguous, input_broadcasted, input_shapes, input_strides, output_strides}; unpackInputsAndApply( [&](auto... Is) { #if defined(ENABLE_HYGON_API) - output[out_idx] = Op{}(typed_inputs[Is.value][indexer(Is.value)]..., args...); + output[out_idx] = Op{}(typedInputPtr(inputs.values[Is.value])[indexer(Is.value)]..., args...); #else - output[out_idx] = Op{}(typed_inputs[Is.value][indexer(Is.value)]..., std::forward(args)...); + output[out_idx] = Op{}(typedInputPtr(inputs.values[Is.value])[indexer(Is.value)]..., std::forward(args)...); #endif }, std::make_index_sequence{}); @@ -300,7 +277,7 @@ INFINIOP_CUDA_KERNEL elementwiseKernel( const ptrdiff_t *__restrict__ output_strides, const ptrdiff_t *__restrict__ input_strides, Tout *output, - const void *const *__restrict__ inputs, + InputPointerArray inputs, size_t offset) { size_t idx = blockIdx.x * blockDim.x + threadIdx.x + offset; @@ -312,7 +289,7 @@ INFINIOP_CUDA_KERNEL elementwiseKernel( unpackInputsAndApply( [&](auto... Is) { output[out_idx] = Op{}.template operator()( - (typedInputPtr(inputs[Is.value])[indexer(Is.value)])...); + (typedInputPtr(inputs.values[Is.value])[indexer(Is.value)])...); }, std::index_sequence_for{}); } @@ -358,9 +335,52 @@ INFINIOP_CUDA_KERNEL inlineMetaElementwiseKernel( struct DeviceImpl::Opaque { std::shared_ptr internal; + void *device_meta = nullptr; + const bool *input_contiguous = nullptr; + const bool *input_broadcasted = nullptr; + const size_t *output_shape = nullptr; + const ptrdiff_t *output_strides = nullptr; + const size_t *input_shapes = nullptr; + const ptrdiff_t *input_strides = nullptr; + infiniStatus_t init_status = INFINI_STATUS_SUCCESS; + + Opaque(const std::shared_ptr &internal_, + const op::elementwise::ElementwiseInfo &info) + : internal(internal_), init_status(initialize(info)) {} + + ~Opaque() { + if (device_meta != nullptr) { + cudaFree(device_meta); + } + } + + infiniStatus_t initialize(const op::elementwise::ElementwiseInfo &info) { + if (info.canUseContiguousFastPath() || info.canUseInlineMetaFastPath()) { + return INFINI_STATUS_SUCCESS; + } + + const auto meta_size = info.getMetaMemSize(); + if (meta_size == 0) { + return INFINI_STATUS_SUCCESS; + } - Opaque(const std::shared_ptr &internal) - : internal(internal) {} + CHECK_CUDA(cudaMalloc(&device_meta, meta_size)); + CHECK_CUDA(cudaMemcpy(device_meta, + info.getMetaStart(), + meta_size, + cudaMemcpyHostToDevice)); + + const auto ndim = info.getNdim(); + const auto input_size = info.getInputSize(); + output_shape = reinterpret_cast(device_meta); + output_strides = reinterpret_cast(output_shape + ndim); + input_shapes = reinterpret_cast(output_strides + ndim); + input_strides = reinterpret_cast(input_shapes + input_size * ndim); + input_contiguous = reinterpret_cast(input_strides + input_size * ndim); + input_broadcasted = input_contiguous + input_size; + + return INFINI_STATUS_SUCCESS; + } /** * @brief Executes an elementwise operation where all inputs and the output share the same data type. @@ -564,73 +584,6 @@ private: return INFINI_STATUS_SUCCESS; } - /** - * @brief Transfers elementwise operation metadata and input pointers from host to device memory. - * - * @tparam N Number of input tensors. - * - * @param info Elementwise operation metadata (shapes, strides, flags, etc.). - * @param workspace Pointer to device workspace memory for storing metadata and input pointers. - * @param h_inputs_arr Host array of input tensor pointers. - * @param d_inputs_arr Input reference to device array of input tensor pointers. - * @param d_input_contiguous Input reference to device array indicating whether each input is contiguous. - * @param d_input_broadcasted Input reference to device array indicating whether each input is broadcasted. - * @param d_output_shape Output reference to device array holding the output tensor shape. - * @param d_output_strides Output reference to device array holding output tensor strides. - * @param d_input_shapes Output reference to flattened input tensor shapes (N * ndim). - * @param d_input_strides Output reference to flattened input tensor strides (N * ndim). - * @param stream CUDA stream used for asynchronous memory transfer. - * @return infiniStatus_t Status indicating success or failure of the memory transfer and setup. - */ - template - infiniStatus_t infoToDevice( - const op::elementwise::ElementwiseInfo &info, - void *workspace, - const void *const *h_inputs_arr, - const void **&d_inputs_arr, - const bool *&d_input_contiguous, - const bool *&d_input_broadcasted, - const size_t *&d_output_shape, - const ptrdiff_t *&d_output_strides, - const size_t *&d_input_shapes, - const ptrdiff_t *&d_input_strides, - cudaStream_t stream) const { - - constexpr auto input_size = N; - const auto ndim = info.getNdim(); - constexpr auto input_arr_size = N * sizeof(*h_inputs_arr); - const int8_t *info_meta_start = info.getMetaStart(); - const int8_t *d_meta_start = reinterpret_cast(workspace) + input_arr_size; - - cudaStreamCaptureStatus capture_status = cudaStreamCaptureStatusNone; - CHECK_CUDA(cudaStreamIsCapturing(stream, &capture_status)); - if (capture_status == cudaStreamCaptureStatusNone) { - CHECK_CUDA(cudaMemcpyAsync(workspace, h_inputs_arr, input_arr_size, cudaMemcpyHostToDevice, stream)); - } else { - // A captured H2D copy from a temporary std::vector would retain an - // invalid host pointer for replay. Kernel arguments are stored by - // value in the graph node instead. - InputPointerArray input_pointers{}; - for (size_t i = 0; i < N; ++i) { - input_pointers.values[i] = h_inputs_arr[i]; - } - storeInputPointers<<<1, N, 0, stream>>>( - reinterpret_cast(workspace), input_pointers); - } - CHECK_CUDA(cudaMemcpyAsync((void *)d_meta_start, info_meta_start, info.getMetaMemSize(), cudaMemcpyHostToDevice, stream)); - - // offset/assign the pointers - d_inputs_arr = reinterpret_cast(workspace); - d_output_shape = reinterpret_cast(d_meta_start); - d_output_strides = reinterpret_cast(d_output_shape + ndim); - d_input_shapes = reinterpret_cast(d_output_strides + ndim); - d_input_strides = reinterpret_cast(d_input_shapes + input_size * ndim); - d_input_contiguous = reinterpret_cast(d_input_strides + input_size * ndim); - d_input_broadcasted = reinterpret_cast(d_input_contiguous + input_size); - - return INFINI_STATUS_SUCCESS; - } - /** * @brief Launches the elementwise kernel for the specified operation. * @@ -664,19 +617,9 @@ private: return INFINI_STATUS_SUCCESS; } - // Device pointers - const void **d_inputs_arr = nullptr; - const bool *d_input_contiguous = nullptr; - const bool *d_input_broadcasted = nullptr; - const size_t *d_output_shape = nullptr; - const ptrdiff_t *d_output_strides = nullptr; - const size_t *d_input_shapes = nullptr; - const ptrdiff_t *d_input_strides = nullptr; - - CHECK_STATUS(infoToDevice(info, workspace, inputs.data(), d_inputs_arr, - d_input_contiguous, d_input_broadcasted, - d_output_shape, d_output_strides, - d_input_shapes, d_input_strides, stream)); + (void)workspace; + InputPointerArray input_ptrs{}; + std::copy_n(inputs.begin(), N, input_ptrs.values); dim3 blockDims(std::min(BLOCK_SIZE, static_cast(internal->maxThreadsPerBlock()))); dim3 gridDims(std::min(uint32_t(CEIL_DIV(output_size, blockDims.x)), static_cast(internal->gridSizeX()))); @@ -685,10 +628,10 @@ private: for (size_t i = 0; i < output_size; i += step) { kernel_func<<>>( output_size, info.getNdim(), info.isOutputContiguous(), - d_input_contiguous, d_input_broadcasted, - d_output_shape, d_input_shapes, - d_output_strides, d_input_strides, - output, reinterpret_cast(d_inputs_arr), + input_contiguous, input_broadcasted, + output_shape, input_shapes, + output_strides, input_strides, + output, input_ptrs, i, std::forward(args)...); } @@ -699,6 +642,9 @@ private: template utils::Result DeviceImpl::create(Args &&...args) { auto opaque = std::make_shared(std::forward(args)...); + if (opaque->init_status != INFINI_STATUS_SUCCESS) { + return opaque->init_status; + } return utils::Result(new DeviceImpl(opaque)); } diff --git a/src/infiniop/elementwise/nvidia/elementwise_nvidia_api.cuh b/src/infiniop/elementwise/nvidia/elementwise_nvidia_api.cuh index 7e7442210..ff60dcac1 100644 --- a/src/infiniop/elementwise/nvidia/elementwise_nvidia_api.cuh +++ b/src/infiniop/elementwise/nvidia/elementwise_nvidia_api.cuh @@ -93,9 +93,9 @@ public: auto info_result = op::elementwise::ElementwiseInfo::create(OUT_DESC, INPUT_DESC_VEC); \ CHECK_RESULT(info_result); \ auto info = info_result.take(); \ - auto workspace_size = info.getMetaMemSize() + info.getInputSize() * sizeof(void *); \ + size_t workspace_size = 0; \ \ - auto device_impl_result = op::elementwise::nvidia::DeviceImpl::create(HANDLE->internal()); \ + auto device_impl_result = op::elementwise::nvidia::DeviceImpl::create(HANDLE->internal(), info); \ CHECK_RESULT(device_impl_result); \ \ *desc_ptr = new Descriptor( \ diff --git a/src/infiniop/ops/hardtanh/nvidia/hardtanh_nvidia.cu b/src/infiniop/ops/hardtanh/nvidia/hardtanh_nvidia.cu index 31ba489ab..d9f728892 100644 --- a/src/infiniop/ops/hardtanh/nvidia/hardtanh_nvidia.cu +++ b/src/infiniop/ops/hardtanh/nvidia/hardtanh_nvidia.cu @@ -87,9 +87,9 @@ infiniStatus_t Descriptor::create( auto info_result = op::elementwise::ElementwiseInfo::create(out_desc, input_desc_vec); CHECK_RESULT(info_result); auto info = info_result.take(); - auto workspace_size = info.getMetaMemSize() + info.getInputSize() * sizeof(void *); + size_t workspace_size = 0; - auto device_impl_result = op::elementwise::nvidia::DeviceImpl::create(handle->internal()); + auto device_impl_result = op::elementwise::nvidia::DeviceImpl::create(handle->internal(), info); CHECK_RESULT(device_impl_result); *desc_ptr = new Descriptor( diff --git a/src/infiniop/ops/mul_scalar/nvidia/mul_scalar_nvidia.cu b/src/infiniop/ops/mul_scalar/nvidia/mul_scalar_nvidia.cu index 06ca80553..3780eff9b 100644 --- a/src/infiniop/ops/mul_scalar/nvidia/mul_scalar_nvidia.cu +++ b/src/infiniop/ops/mul_scalar/nvidia/mul_scalar_nvidia.cu @@ -75,9 +75,7 @@ infiniStatus_t calculateMulScalar( return launchMulScalarKernel(info.numel(), output, input, alpha, stream); } - if (workspace_size < info.elementwise_info.getMetaMemSize() + sizeof(void *)) { - return INFINI_STATUS_INSUFFICIENT_WORKSPACE; - } + (void)workspace_size; return device_info->calculate<256, MulScalarOp, T>( info.elementwise_info, @@ -105,8 +103,8 @@ infiniStatus_t Descriptor::create( CHECK_RESULT(result); auto info = result.take(); - auto workspace_size = info.elementwise_info.getMetaMemSize() + sizeof(void *); - auto device_impl_result = op::elementwise::nvidia::DeviceImpl::create(handle->internal()); + size_t workspace_size = 0; + auto device_impl_result = op::elementwise::nvidia::DeviceImpl::create(handle->internal(), info.elementwise_info); CHECK_RESULT(device_impl_result); *desc_ptr = new Descriptor(