diff --git a/src/infiniop/elementwise/nvidia/elementwise_nvidia.cuh b/src/infiniop/elementwise/nvidia/elementwise_nvidia.cuh index f8cb457ec..d93abf06e 100644 --- a/src/infiniop/elementwise/nvidia/elementwise_nvidia.cuh +++ b/src/infiniop/elementwise/nvidia/elementwise_nvidia.cuh @@ -172,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{}); @@ -278,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; @@ -290,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{}); } @@ -336,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); + } + } - Opaque(const std::shared_ptr &internal) - : internal(internal) {} + 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; + } + + 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. @@ -542,60 +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; - - // copy the input pointer array and meta to device - CHECK_CUDA(cudaMemcpyAsync(workspace, h_inputs_arr, input_arr_size, cudaMemcpyHostToDevice, stream)); - 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. * @@ -629,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()))); @@ -650,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)...); } @@ -664,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/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/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/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, 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(