Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
148 changes: 78 additions & 70 deletions src/infiniop/elementwise/metax/elementwise_metax.h
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,11 @@ __device__ __forceinline__ const T *typedInputPtr(const void *ptr) {
return reinterpret_cast<const T *>(ptr);
}

template <size_t N>
struct InputPointerArray {
const void *values[N];
};

// Generic aligned N-element pack used for vectorized load/store.
template <typename T, int N>
struct alignas(sizeof(T) * N) Pack {
Expand Down Expand Up @@ -124,10 +129,15 @@ template <size_t N, typename Op, typename Tdata, typename VecT, int V, typename.
INFINIOP_METAX_KERNEL elementwiseVecKernel(
size_t output_size,
Tdata *__restrict__ output,
const void *const *__restrict__ inputs,
InputPointerArray<N> inputs,
Args... args) {

const Tdata *const *typed_inputs = reinterpret_cast<const Tdata *const *>(inputs);
const Tdata *typed_inputs[N];
#pragma unroll
for (size_t i = 0; i < N; ++i) {
typed_inputs[i] = typedInputPtr<Tdata>(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;
Expand Down Expand Up @@ -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<N> 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<const Tdata *const *>(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>(args)...);
output[out_idx] = Op{}(
typedInputPtr<Tdata>(inputs.values[Is.value])[indexer(Is.value)]...,
std::forward<Args>(args)...);
},
std::make_index_sequence<N>{});
}
Expand All @@ -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<sizeof...(Tin)> inputs,
size_t offset) {

size_t idx = blockIdx.x * blockDim.x + threadIdx.x + offset;
Expand All @@ -216,17 +227,56 @@ INFINIOP_METAX_KERNEL elementwiseKernel(
unpackInputsAndApply(
[&](auto... Is) {
output[out_idx] = Op{}.template operator()<Tout, Tin...>(
(typedInputPtr<Tin>(inputs[Is.value])[indexer(Is.value)])...);
(typedInputPtr<Tin>(inputs.values[Is.value])[indexer(Is.value)])...);
},
std::index_sequence_for<Tin...>{});
}
}

struct DeviceImpl::Opaque {
std::shared_ptr<device::metax::Handle::Internal> 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<device::metax::Handle::Internal> &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<device::metax::Handle::Internal> &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<const size_t *>(device_meta);
output_strides = reinterpret_cast<const ptrdiff_t *>(output_shape + ndim);
input_shapes = reinterpret_cast<const size_t *>(output_strides + ndim);
input_strides = reinterpret_cast<const ptrdiff_t *>(input_shapes + input_size * ndim);
input_contiguous = reinterpret_cast<const bool *>(input_strides + input_size * ndim);
input_broadcasted = input_contiguous + input_size;

return INFINI_STATUS_SUCCESS;
}

template <uint32_t BLOCK_SIZE, size_t N, typename Op, typename Tdata, typename... Args>
infiniStatus_t calculateImpl(const op::elementwise::ElementwiseInfo &info,
Expand All @@ -239,18 +289,19 @@ struct DeviceImpl::Opaque {
if (canUseVecPath<Tdata, N>(info, output, inputs)) {
return launchElementwiseVecKernel<BLOCK_SIZE, N, Op, Tdata,
typename VecInfo<Tdata>::Type,
VecInfo<Tdata>::pack_size>(
VecInfo<Tdata>::pack_size,
std::decay_t<Args>...>(
info, workspace,
reinterpret_cast<Tdata *>(output), inputs, stream,
std::forward<Args>(args)...);
std::decay_t<Args>(args)...);
}
}
return launchElementwiseKernel<BLOCK_SIZE, N>(
info, workspace,
reinterpret_cast<Tdata *>(output), inputs,
elementwiseKernel<N, Op, Tdata, Args...>,
elementwiseKernel<N, Op, Tdata, std::decay_t<Args>...>,
stream,
std::forward<Args>(args)...);
std::decay_t<Args>(args)...);
}

template <uint32_t BLOCK_SIZE, size_t N, typename Op, typename Tout, typename... Tin, typename... Args,
Expand Down Expand Up @@ -283,9 +334,9 @@ struct DeviceImpl::Opaque {
return INFINI_STATUS_SUCCESS;
}

CHECK_METAX(hcMemcpyAsync(workspace, inputs.data(), N * sizeof(*inputs.data()),
hcMemcpyHostToDevice, stream));
const void **d_inputs_arr = reinterpret_cast<const void **>(workspace);
(void)workspace;
InputPointerArray<N> input_ptrs{};
std::copy_n(inputs.begin(), N, input_ptrs.values);

dim3 blockDims(std::min(BLOCK_SIZE, static_cast<uint32_t>(internal->maxThreadsPerBlock())));
const size_t num_packs = output_size / V;
Expand All @@ -299,43 +350,7 @@ struct DeviceImpl::Opaque {

elementwiseVecKernel<N, Op, Tdata, VecT, V, Args...>
<<<gridDims, blockDims, 0, stream>>>(
output_size, output, d_inputs_arr, std::forward<Args>(args)...);

return INFINI_STATUS_SUCCESS;
}

template <size_t N>
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<int8_t *>(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<const void **>(workspace);
d_output_shape = reinterpret_cast<const size_t *>(d_meta_start);
d_output_strides = reinterpret_cast<const ptrdiff_t *>(d_output_shape + ndim);
d_input_shapes = reinterpret_cast<const size_t *>(d_output_strides + ndim);
d_input_strides = reinterpret_cast<const ptrdiff_t *>(d_input_shapes + input_size * ndim);
d_input_contiguous = reinterpret_cast<const bool *>(d_input_strides + input_size * ndim);
d_input_broadcasted = reinterpret_cast<const bool *>(d_input_contiguous + input_size);
output_size, output, input_ptrs, std::forward<Args>(args)...);

return INFINI_STATUS_SUCCESS;
}
Expand All @@ -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<N>(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<N> input_ptrs{};
std::copy_n(inputs.begin(), N, input_ptrs.values);

dim3 blockDims(std::min(BLOCK_SIZE, static_cast<uint32_t>(internal->maxThreadsPerBlock())));
dim3 gridDims(std::min(uint32_t(CEIL_DIV(output_size, blockDims.x)), static_cast<uint32_t>(internal->gridSizeX())));
Expand All @@ -376,10 +381,10 @@ struct DeviceImpl::Opaque {
for (size_t i = 0; i < output_size; i += step) {
kernel_func<<<gridDims, blockDims, 0, stream>>>(
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<const void **>(d_inputs_arr),
input_contiguous, input_broadcasted,
output_shape, input_shapes,
output_strides, input_strides,
output, input_ptrs,
i, std::forward<Args>(args)...);
}

Expand All @@ -390,6 +395,9 @@ struct DeviceImpl::Opaque {
template <typename... Args>
utils::Result<DeviceImpl *> DeviceImpl::create(Args &&...args) {
auto opaque = std::make_shared<Opaque>(std::forward<Args>(args)...);
if (opaque->init_status != INFINI_STATUS_SUCCESS) {
return opaque->init_status;
}
return utils::Result<DeviceImpl *>(new DeviceImpl(opaque));
}

Expand Down
32 changes: 16 additions & 16 deletions src/infiniop/elementwise/metax/elementwise_metax_api.h
Original file line number Diff line number Diff line change
Expand Up @@ -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__
4 changes: 2 additions & 2 deletions src/infiniop/ops/hardtanh/metax/hardtanh_metax.maca
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
8 changes: 3 additions & 5 deletions src/infiniop/ops/mul_scalar/metax/mul_scalar_metax.maca
Original file line number Diff line number Diff line change
Expand Up @@ -81,9 +81,7 @@ infiniStatus_t calculateMulScalar(
return launchMulScalarKernel<T>(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,
Expand Down Expand Up @@ -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(
Expand Down
Loading