66#include " ../../devices/moore/moore_kernel_common.h"
77#include " elementwise_moore_api.h"
88
9+ #include < algorithm>
10+
911namespace op ::elementwise::moore {
1012template <typename T>
1113__device__ __forceinline__ const T *typedInputPtr (const void *ptr) {
1214 return reinterpret_cast <const T *>(ptr);
1315}
1416
17+ template <size_t N>
18+ struct InputPointerArray {
19+ const void *values[N];
20+ };
21+
1522__device__ __forceinline__ size_t getOutputIndex (size_t idx, bool is_contiguous, size_t ndim,
1623 const size_t *shape, const ptrdiff_t *strides) {
1724 return is_contiguous ? idx : device::moore::indexToOffset (idx, ndim, shape, strides);
@@ -50,20 +57,20 @@ INFINIOP_MOORE_KERNEL elementwiseKernel(
5057 const ptrdiff_t *__restrict__ output_strides,
5158 const ptrdiff_t *__restrict__ input_strides,
5259 Tdata *output,
53- const void * const * inputs,
60+ InputPointerArray<N> inputs,
5461 size_t offset,
5562 Args... args) {
5663
5764 size_t idx = blockIdx.x * blockDim.x + threadIdx.x + offset;
5865
5966 if (idx < output_size) {
60- const Tdata *const *typed_inputs = reinterpret_cast <const Tdata *const *>(inputs);
6167 size_t out_idx = getOutputIndex (idx, output_contiguous, ndim, output_shape, output_strides);
6268 InputIndexer indexer{idx, ndim, input_contiguous, input_broadcasted, input_shapes, input_strides, output_strides};
6369
6470 unpackInputsAndApply (
6571 [&](auto ... Is) {
66- output[out_idx] = Op{}(typed_inputs[Is.value ][indexer (Is.value )]..., std::forward<Args>(args)...);
72+ output[out_idx] = Op{}(
73+ typedInputPtr<Tdata>(inputs.values [Is.value ])[indexer (Is.value )]..., std::forward<Args>(args)...);
6774 },
6875 std::make_index_sequence<N>{});
6976 }
@@ -81,7 +88,7 @@ INFINIOP_MOORE_KERNEL elementwiseKernel(
8188 const ptrdiff_t *__restrict__ output_strides,
8289 const ptrdiff_t *__restrict__ input_strides,
8390 Tout *output,
84- const void * const *__restrict__ inputs,
91+ InputPointerArray< sizeof ...(Tin)> inputs,
8592 size_t offset) {
8693
8794 size_t idx = blockIdx.x * blockDim.x + threadIdx.x + offset;
@@ -93,17 +100,56 @@ INFINIOP_MOORE_KERNEL elementwiseKernel(
93100 unpackInputsAndApply (
94101 [&](auto ... Is) {
95102 output[out_idx] = Op{}.template operator ()<Tout, Tin...>(
96- (typedInputPtr<Tin>(inputs[Is.value ])[indexer (Is.value )])...);
103+ (typedInputPtr<Tin>(inputs. values [Is.value ])[indexer (Is.value )])...);
97104 },
98105 std::index_sequence_for<Tin...>{});
99106 }
100107}
101108
102109struct DeviceImpl ::Opaque {
103110 std::shared_ptr<device::moore::Handle::Internal> internal;
111+ void *device_meta = nullptr ;
112+ const bool *input_contiguous = nullptr ;
113+ const bool *input_broadcasted = nullptr ;
114+ const size_t *output_shape = nullptr ;
115+ const ptrdiff_t *output_strides = nullptr ;
116+ const size_t *input_shapes = nullptr ;
117+ const ptrdiff_t *input_strides = nullptr ;
118+ infiniStatus_t init_status = INFINI_STATUS_SUCCESS ;
119+
120+ Opaque (const std::shared_ptr<device::moore::Handle::Internal> &internal_,
121+ const op::elementwise::ElementwiseInfo &info)
122+ : internal(internal_), init_status(initialize(info)) {}
123+
124+ ~Opaque () {
125+ if (device_meta != nullptr ) {
126+ musaFree (device_meta);
127+ }
128+ }
129+
130+ infiniStatus_t initialize (const op::elementwise::ElementwiseInfo &info) {
131+ const auto meta_size = info.getMetaMemSize ();
132+ if (meta_size == 0 ) {
133+ return INFINI_STATUS_SUCCESS ;
134+ }
135+
136+ CHECK_MOORE (musaMalloc (&device_meta, meta_size));
137+ CHECK_MOORE (musaMemcpy (device_meta,
138+ info.getMetaStart (),
139+ meta_size,
140+ musaMemcpyHostToDevice));
141+
142+ const auto ndim = info.getNdim ();
143+ const auto input_size = info.getInputSize ();
144+ output_shape = reinterpret_cast <const size_t *>(device_meta);
145+ output_strides = reinterpret_cast <const ptrdiff_t *>(output_shape + ndim);
146+ input_shapes = reinterpret_cast <const size_t *>(output_strides + ndim);
147+ input_strides = reinterpret_cast <const ptrdiff_t *>(input_shapes + input_size * ndim);
148+ input_contiguous = reinterpret_cast <const bool *>(input_strides + input_size * ndim);
149+ input_broadcasted = input_contiguous + input_size;
104150
105- Opaque ( const std::shared_ptr<device::moore::Handle::Internal> &internal)
106- : internal(internal) { }
151+ return INFINI_STATUS_SUCCESS ;
152+ }
107153
108154 template <uint32_t BLOCK_SIZE , size_t N, typename Op, typename Tdata, typename ... Args>
109155 infiniStatus_t calculateImpl (const op::elementwise::ElementwiseInfo &info,
@@ -136,42 +182,6 @@ struct DeviceImpl::Opaque {
136182 }
137183
138184private:
139- template <size_t N>
140- infiniStatus_t infoToDevice (
141- const op::elementwise::ElementwiseInfo &info,
142- void *workspace,
143- const void *const *h_inputs_arr,
144- const void **&d_inputs_arr,
145- const bool *&d_input_contiguous,
146- const bool *&d_input_broadcasted,
147- const size_t *&d_output_shape,
148- const ptrdiff_t *&d_output_strides,
149- const size_t *&d_input_shapes,
150- const ptrdiff_t *&d_input_strides,
151- musaStream_t stream) const {
152-
153- constexpr auto input_size = N;
154- const auto ndim = info.getNdim ();
155- constexpr auto input_arr_size = N * sizeof (*h_inputs_arr);
156- const int8_t *info_meta_start = info.getMetaStart ();
157- const int8_t *d_meta_start = reinterpret_cast <int8_t *>(workspace) + input_arr_size;
158-
159- // copy the input pointer array and meta to device
160- CHECK_MOORE (musaMemcpyAsync (workspace, h_inputs_arr, input_arr_size, musaMemcpyHostToDevice, stream));
161- CHECK_MOORE (musaMemcpyAsync ((void *)d_meta_start, info_meta_start, info.getMetaMemSize (), musaMemcpyHostToDevice, stream));
162-
163- // offset/assign the pointers
164- d_inputs_arr = reinterpret_cast <const void **>(workspace);
165- d_output_shape = reinterpret_cast <const size_t *>(d_meta_start);
166- d_output_strides = reinterpret_cast <const ptrdiff_t *>(d_output_shape + ndim);
167- d_input_shapes = reinterpret_cast <const size_t *>(d_output_strides + ndim);
168- d_input_strides = reinterpret_cast <const ptrdiff_t *>(d_input_shapes + input_size * ndim);
169- d_input_contiguous = reinterpret_cast <const bool *>(d_input_strides + input_size * ndim);
170- d_input_broadcasted = reinterpret_cast <const bool *>(d_input_contiguous + input_size);
171-
172- return INFINI_STATUS_SUCCESS ;
173- }
174-
175185 template <uint32_t BLOCK_SIZE , size_t N, typename KernelFunc, typename Tout, typename ... Args>
176186 infiniStatus_t launchElementwiseKernel (
177187 const op::elementwise::ElementwiseInfo &info,
@@ -187,19 +197,10 @@ struct DeviceImpl::Opaque {
187197 return INFINI_STATUS_SUCCESS ;
188198 }
189199
190- // Device pointers
191- const void **d_inputs_arr = nullptr ;
192- const bool *d_input_contiguous = nullptr ;
193- const bool *d_input_broadcasted = nullptr ;
194- const size_t *d_output_shape = nullptr ;
195- const ptrdiff_t *d_output_strides = nullptr ;
196- const size_t *d_input_shapes = nullptr ;
197- const ptrdiff_t *d_input_strides = nullptr ;
198-
199- CHECK_STATUS (infoToDevice<N>(info, workspace, inputs.data (), d_inputs_arr,
200- d_input_contiguous, d_input_broadcasted,
201- d_output_shape, d_output_strides,
202- d_input_shapes, d_input_strides, stream));
200+ (void )workspace;
201+ CHECK_OR_RETURN (inputs.size () == N, INFINI_STATUS_BAD_PARAM );
202+ InputPointerArray<N> input_ptrs{};
203+ std::copy_n (inputs.begin (), N, input_ptrs.values );
203204
204205 dim3 blockDims (std::min (BLOCK_SIZE , static_cast <uint32_t >(internal->maxThreadsPerBlock ())));
205206 dim3 gridDims (std::min (uint32_t (CEIL_DIV (output_size, blockDims.x )), static_cast <uint32_t >(internal->gridSizeX ())));
@@ -208,10 +209,10 @@ struct DeviceImpl::Opaque {
208209 for (size_t i = 0 ; i < output_size; i += step) {
209210 kernel_func<<<gridDims, blockDims, 0 , stream>>>(
210211 output_size, info.getNdim (), info.isOutputContiguous (),
211- d_input_contiguous, d_input_broadcasted ,
212- d_output_shape, d_input_shapes ,
213- d_output_strides, d_input_strides ,
214- output, reinterpret_cast < const void **>(d_inputs_arr) ,
212+ input_contiguous, input_broadcasted ,
213+ output_shape, input_shapes ,
214+ output_strides, input_strides ,
215+ output, input_ptrs ,
215216 i, std::forward<Args>(args)...);
216217 }
217218
@@ -222,6 +223,9 @@ struct DeviceImpl::Opaque {
222223template <typename ... Args>
223224utils::Result<DeviceImpl *> DeviceImpl::create (Args &&...args) {
224225 auto opaque = std::make_shared<Opaque>(std::forward<Args>(args)...);
226+ if (opaque->init_status != INFINI_STATUS_SUCCESS ) {
227+ return opaque->init_status ;
228+ }
225229 return utils::Result<DeviceImpl *>(new DeviceImpl (opaque));
226230}
227231
0 commit comments