From 9a6feff8bdd33ea57f712b5c81ee29781b75d771 Mon Sep 17 00:00:00 2001 From: wooway777 Date: Fri, 4 Sep 2026 08:03:18 +0000 Subject: [PATCH] feat: enable graph on metax elementwise --- .../elementwise/metax/elementwise_metax.h | 148 +++++++++--------- .../elementwise/metax/elementwise_metax_api.h | 32 ++-- .../ops/hardtanh/metax/hardtanh_metax.maca | 4 +- .../mul_scalar/metax/mul_scalar_metax.maca | 8 +- 4 files changed, 99 insertions(+), 93 deletions(-) diff --git a/src/infiniop/elementwise/metax/elementwise_metax.h b/src/infiniop/elementwise/metax/elementwise_metax.h index 52e4ddf54..35e2aeb6f 100644 --- a/src/infiniop/elementwise/metax/elementwise_metax.h +++ b/src/infiniop/elementwise/metax/elementwise_metax.h @@ -13,6 +13,11 @@ __device__ __forceinline__ const T *typedInputPtr(const void *ptr) { return reinterpret_cast(ptr); } +template +struct InputPointerArray { + const void *values[N]; +}; + // Generic aligned N-element pack used for vectorized load/store. template struct alignas(sizeof(T) * N) Pack { @@ -124,10 +129,15 @@ template inputs, Args... args) { - const Tdata *const *typed_inputs = reinterpret_cast(inputs); + const Tdata *typed_inputs[N]; +#pragma unroll + for (size_t i = 0; i < N; ++i) { + typed_inputs[i] = typedInputPtr(inputs.values[i]); + } + const size_t num_packs = output_size / V; const size_t tail_start = num_packs * V; const size_t tid = blockIdx.x * blockDim.x + threadIdx.x; @@ -173,20 +183,21 @@ INFINIOP_METAX_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) { - 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)...); }, std::make_index_sequence{}); } @@ -204,7 +215,7 @@ INFINIOP_METAX_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; @@ -216,7 +227,7 @@ INFINIOP_METAX_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{}); } @@ -224,9 +235,48 @@ INFINIOP_METAX_KERNEL elementwiseKernel( 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) { + hcFree(device_meta); + } + } - Opaque(const std::shared_ptr &internal) - : internal(internal) {} + infiniStatus_t initialize(const op::elementwise::ElementwiseInfo &info) { + const auto meta_size = info.getMetaMemSize(); + if (meta_size == 0) { + return INFINI_STATUS_SUCCESS; + } + + CHECK_METAX(hcMalloc(&device_meta, meta_size)); + CHECK_METAX(hcMemcpy(device_meta, + info.getMetaStart(), + meta_size, + hcMemcpyHostToDevice)); + + 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; + } template infiniStatus_t calculateImpl(const op::elementwise::ElementwiseInfo &info, @@ -239,18 +289,19 @@ struct DeviceImpl::Opaque { if (canUseVecPath(info, output, inputs)) { return launchElementwiseVecKernel::Type, - VecInfo::pack_size>( + VecInfo::pack_size, + std::decay_t...>( info, workspace, reinterpret_cast(output), inputs, stream, - std::forward(args)...); + std::decay_t(args)...); } } return launchElementwiseKernel( info, workspace, reinterpret_cast(output), inputs, - elementwiseKernel, + elementwiseKernel...>, stream, - std::forward(args)...); + std::decay_t(args)...); } template (workspace); + (void)workspace; + InputPointerArray input_ptrs{}; + std::copy_n(inputs.begin(), N, input_ptrs.values); dim3 blockDims(std::min(BLOCK_SIZE, static_cast(internal->maxThreadsPerBlock()))); const size_t num_packs = output_size / V; @@ -299,43 +350,7 @@ struct DeviceImpl::Opaque { elementwiseVecKernel <<>>( - output_size, output, d_inputs_arr, std::forward(args)...); - - return INFINI_STATUS_SUCCESS; - } - - 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, - hcStream_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_METAX(hcMemcpyAsync(workspace, h_inputs_arr, input_arr_size, hcMemcpyHostToDevice, stream)); - CHECK_METAX(hcMemcpyAsync((void *)d_meta_start, info_meta_start, info.getMetaMemSize(), hcMemcpyHostToDevice, 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); + output_size, output, input_ptrs, std::forward(args)...); return INFINI_STATUS_SUCCESS; } @@ -355,19 +370,9 @@ struct DeviceImpl::Opaque { 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()))); @@ -376,10 +381,10 @@ struct DeviceImpl::Opaque { 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)...); } @@ -390,6 +395,9 @@ struct DeviceImpl::Opaque { 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/metax/elementwise_metax_api.h b/src/infiniop/elementwise/metax/elementwise_metax_api.h index b59c14da5..b458c49b2 100644 --- a/src/infiniop/elementwise/metax/elementwise_metax_api.h +++ b/src/infiniop/elementwise/metax/elementwise_metax_api.h @@ -38,22 +38,22 @@ class DeviceImpl final { Args &&...args); }; } // namespace op::elementwise::metax -#define CREATE_ELEMENTWISE_METAX_DESCRIPTOR(HANDLE, DTYPE, OUT_DESC, INPUT_DESC_VEC) \ - \ - 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 *); \ - \ - auto device_impl_result = op::elementwise::metax::DeviceImpl::create(HANDLE->internal()); \ - CHECK_RESULT(device_impl_result); \ - \ - *desc_ptr = new Descriptor( \ - DTYPE, \ - std::move(info), \ - std::move(device_impl_result.take()), \ - workspace_size, \ - HANDLE->device, \ +#define CREATE_ELEMENTWISE_METAX_DESCRIPTOR(HANDLE, DTYPE, OUT_DESC, INPUT_DESC_VEC) \ + \ + auto info_result = op::elementwise::ElementwiseInfo::create(OUT_DESC, INPUT_DESC_VEC); \ + CHECK_RESULT(info_result); \ + auto info = info_result.take(); \ + size_t workspace_size = 0; \ + \ + auto device_impl_result = op::elementwise::metax::DeviceImpl::create(HANDLE->internal(), info); \ + CHECK_RESULT(device_impl_result); \ + \ + *desc_ptr = new Descriptor( \ + DTYPE, \ + std::move(info), \ + std::move(device_impl_result.take()), \ + workspace_size, \ + HANDLE->device, \ HANDLE->device_id); #endif // __INFINIOP_ELEMENTWISE_METAX_API_H__ diff --git a/src/infiniop/ops/hardtanh/metax/hardtanh_metax.maca b/src/infiniop/ops/hardtanh/metax/hardtanh_metax.maca index 596316e23..ca11b161b 100644 --- a/src/infiniop/ops/hardtanh/metax/hardtanh_metax.maca +++ b/src/infiniop/ops/hardtanh/metax/hardtanh_metax.maca @@ -45,9 +45,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::metax::DeviceImpl::create(handle->internal()); + auto device_impl_result = op::elementwise::metax::DeviceImpl::create(handle->internal(), info); CHECK_RESULT(device_impl_result); *desc_ptr = new Descriptor( diff --git a/src/infiniop/ops/mul_scalar/metax/mul_scalar_metax.maca b/src/infiniop/ops/mul_scalar/metax/mul_scalar_metax.maca index b97a9e72f..99c755973 100644 --- a/src/infiniop/ops/mul_scalar/metax/mul_scalar_metax.maca +++ b/src/infiniop/ops/mul_scalar/metax/mul_scalar_metax.maca @@ -81,9 +81,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, @@ -111,8 +109,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::metax::DeviceImpl::create(handle->internal()); + size_t workspace_size = 0; + auto device_impl_result = op::elementwise::metax::DeviceImpl::create(handle->internal(), info.elementwise_info); CHECK_RESULT(device_impl_result); *desc_ptr = new Descriptor(