Skip to content
Open
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
147 changes: 147 additions & 0 deletions src/native/cambricon/ops/gelu/kernel.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,147 @@
#ifndef INFINI_OPS_CAMBRICON_GELU_CNNL_H_
#define INFINI_OPS_CAMBRICON_GELU_CNNL_H_

#include <algorithm>
#include <cassert>

#include "base/gelu.h"
#include "native/cambricon/cnnl_utils.h"
#include "native/cambricon/cnrt_utils.h"
#include "native/cambricon/data_type_.h"

namespace infini::ops {

template <typename T>
void GeluUnion(void* workspace, int core_per_cluster, int cluster_count,
cnrtQueue_t queue, const void* input, void* out,
const size_t* input_shape, const ptrdiff_t* input_strides,
const size_t* out_shape, const ptrdiff_t* out_strides,
size_t output_size, int ndim, bool approximate,
bool input_contiguous, bool out_contiguous);

template <>
class Operator<Gelu, Device::Type::kCambricon> : public Gelu {
public:
Operator(const Tensor input, const std::string approximate, Tensor out)
: Gelu{input, approximate, out} {
assert(input_shape_ == out_shape_ &&
"`CambriconGelu` requires matching input and output shapes.");
assert(input_type_ == out_type_ &&
"`CambriconGelu` requires matching input and output dtypes.");
assert(input.device() == out.device() &&
"`CambriconGelu` requires input and output on the same device.");
assert((input_type_ == DataType::kFloat16 ||
input_type_ == DataType::kBFloat16 ||
input_type_ == DataType::kFloat32) &&
"`CambriconGelu` supports float16, bfloat16, and float32 only.");
assert(!out.HasBroadcastDim() &&
"`CambriconGelu` output must not have broadcast dimensions.");
assert(std::all_of(input_strides_.begin(), input_strides_.end(),
[](auto stride) { return stride >= 0; }) &&
"`CambriconGelu` does not support negative input strides.");
assert(std::all_of(out_strides_.begin(), out_strides_.end(),
[](auto stride) { return stride >= 0; }) &&
"`CambriconGelu` does not support negative output strides.");

if (output_size_ == 0) {
return;
}

if (approximate_ == "tanh" || input_type_ == DataType::kBFloat16) {
cnrt_utils::GetLaunchConfig(input.device(), &core_per_cluster_,
&cluster_count_);
const auto workspace_size = workspace_size_in_bytes();
if (workspace_size != 0) {
CNRT_CHECK(cnrtMalloc(&default_workspace_, workspace_size));
}
return;
}

cnnl_handle_ = cnnl_utils::CreateHandle();
input_desc_ = cnnl_utils::MakeTensorDescriptor(input_type_, input_shape_,
input_strides_);
out_desc_ =
cnnl_utils::MakeTensorDescriptor(out_type_, out_shape_, out_strides_);

INFINI_OPS_CNNL_CHECK(cnnlCreateActivationDescriptor(&activation_desc_));
const cnnlActivationMode_t mode = CNNL_ACTIVATION_GELU;
const cnnlComputationPreference_t preference =
CNNL_COMPUTATION_HIGH_PRECISION;
const cnnlNanPropagation_t nan_propagation = CNNL_PROPAGATE_NAN;
const bool use_approximation = false;
INFINI_OPS_CNNL_CHECK(cnnlSetActivationDescAttr(
activation_desc_, CNNL_ACTIVATION_MODE, &mode, sizeof(mode)));
INFINI_OPS_CNNL_CHECK(
cnnlSetActivationDescAttr(activation_desc_, CNNL_ACTIVATION_PREFERENCE,
&preference, sizeof(preference)));
INFINI_OPS_CNNL_CHECK(
cnnlSetActivationDescAttr(activation_desc_, CNNL_ACTIVATION_NAN_PROP,
&nan_propagation, sizeof(nan_propagation)));
INFINI_OPS_CNNL_CHECK(cnnlSetActivationDescAttr(
activation_desc_, CNNL_ACTIVATION_APPROXIMATE, &use_approximation,
sizeof(use_approximation)));
}

~Operator() {
if (default_workspace_) {
(void)cnrtFree(default_workspace_);
}
if (activation_desc_) {
(void)cnnlDestroyActivationDescriptor(activation_desc_);
}
}

void operator()(const Tensor input, const std::string approximate,
Tensor out) const override {
assert(approximate == approximate_ &&
"`CambriconGelu` attributes changed after descriptor creation.");
if (output_size_ == 0) {
return;
}

if (approximate_ == "tanh" || input_type_ == DataType::kBFloat16) {
auto queue = static_cast<cnrtQueue_t>(stream_ ? stream_ : 0);
auto* workspace = workspace_ ? workspace_ : default_workspace_;
DispatchFunc<
Device::Type::kCambricon,
List<DataType::kFloat16, DataType::kBFloat16, DataType::kFloat32>>(
{out_type_},
[&](auto tag) {
using T = typename decltype(tag)::type;
GeluUnion<T>(workspace, core_per_cluster_, cluster_count_, queue,
input.data(), out.data(), input_shape_.data(),
input_strides_.data(), out_shape_.data(),
out_strides_.data(), output_size_,
static_cast<int>(ndim_), approximate_ == "tanh",
is_input_contiguous_, is_out_contiguous_);
},
"CambriconGelu::operator() - output dispatch");
return;
}

INFINI_OPS_CNNL_CHECK(cnnlSetQueue(
cnnl_handle_.get(), static_cast<cnrtQueue_t>(stream_ ? stream_ : 0)));
INFINI_OPS_CNNL_CHECK(cnnlActivationForward(
cnnl_handle_.get(), activation_desc_, nullptr, input_desc_.get(),
input.data(), nullptr, out_desc_.get(), out.data()));
}

std::size_t workspace_size_in_bytes() const override {
return (approximate_ == "tanh" || input_type_ == DataType::kBFloat16)
? ndim_ * (2 * sizeof(size_t) + 2 * sizeof(ptrdiff_t))
: 0;
}

private:
cnnl_utils::Handle cnnl_handle_{};
cnnl_utils::TensorDescriptor input_desc_{};
cnnl_utils::TensorDescriptor out_desc_{};
cnnlActivationDescriptor_t activation_desc_{nullptr};
void* default_workspace_{nullptr};
int core_per_cluster_{0};
int cluster_count_{0};
};

} // namespace infini::ops

#endif
172 changes: 172 additions & 0 deletions src/native/cambricon/ops/gelu/kernel.mlu
Original file line number Diff line number Diff line change
@@ -0,0 +1,172 @@
#include <cmath>
#include <cstddef>
#include <type_traits>

#include "kernel.h"

namespace infini::ops {
namespace {

__nram__ char gelu_nram_buffer[NRAM_MAX_SIZE] __attribute__((aligned(128)));

__mlu_device__ ptrdiff_t LogicalToOffset(size_t logical_index, int ndim,
const size_t* shape,
const ptrdiff_t* strides) {
ptrdiff_t offset = 0;
for (int dim = ndim - 1; dim >= 0; --dim) {
const size_t coordinate = logical_index % shape[dim];
logical_index /= shape[dim];
offset += static_cast<ptrdiff_t>(coordinate) * strides[dim];
}
return offset;
}

template <typename T>
__mlu_device__ void ComputeGelu(const T* input, T* output, float* values,
size_t count, bool approximate) {
if constexpr (std::is_same_v<T, __half>) {
__bang_half2float(values, reinterpret_cast<half*>(const_cast<T*>(input)),
count);
} else if constexpr (std::is_same_v<T, __bang_bfloat16>) {
__bang_bfloat162float(values, const_cast<T*>(input), count);
} else {
__memcpy(values, const_cast<T*>(input), count * sizeof(float), NRAM2NRAM);
}

constexpr float kSqrtTwoOverPi = 0.7978845608028654F;
constexpr float kCubicCoefficient = 0.044715F;
for (size_t i = 0; i < count; ++i) {
const float x = values[i];
if (approximate) {
const float inner = kSqrtTwoOverPi * (x + kCubicCoefficient * x * x * x);
values[i] = 0.5F * x * (1.0F + tanhf(inner));
} else {
constexpr float kInverseSqrtTwo = 0.7071067811865475F;
values[i] = 0.5F * x * (1.0F + erff(x * kInverseSqrtTwo));
}
}

if constexpr (std::is_same_v<T, __half>) {
__bang_float2half(reinterpret_cast<half*>(output), values, count);
} else if constexpr (std::is_same_v<T, __bang_bfloat16>) {
__bang_float2bfloat16(output, values, count);
} else {
__memcpy(output, values, count * sizeof(float), NRAM2NRAM);
}
}

template <typename T>
__mlu_global__ void GeluKernel(const T* input, T* output,
const size_t* input_shape,
const ptrdiff_t* input_strides,
const size_t* out_shape,
const ptrdiff_t* out_strides, size_t output_size,
int ndim, bool approximate,
bool input_contiguous, bool out_contiguous) {
const size_t elements_per_task = (output_size + taskDim - 1) / taskDim;
const size_t begin = taskId * elements_per_task;
const size_t end = begin + elements_per_task < output_size
? begin + elements_per_task
: output_size;
if (begin >= end) {
return;
}

size_t block_size = NRAM_MAX_SIZE / (2 * sizeof(T) + sizeof(float));
block_size = block_size / 64 * 64;
if (block_size == 0) {
block_size = 1;
}
auto* input_buffer = reinterpret_cast<T*>(gelu_nram_buffer);
auto* output_buffer = input_buffer + block_size;
auto* values = reinterpret_cast<float*>(output_buffer + block_size);

size_t processed = begin;
while (processed < end) {
const size_t current =
block_size < end - processed ? block_size : end - processed;
if (input_contiguous) {
__memcpy(input_buffer, input + processed, current * sizeof(T),
GDRAM2NRAM);
} else {
for (size_t i = 0; i < current; ++i) {
const size_t logical = processed + i;
input_buffer[i] =
input[LogicalToOffset(logical, ndim, input_shape, input_strides)];
}
}

ComputeGelu(input_buffer, output_buffer, values, current, approximate);
if (out_contiguous) {
__memcpy(output + processed, output_buffer, current * sizeof(T),
NRAM2GDRAM);
} else {
for (size_t i = 0; i < current; ++i) {
const size_t logical = processed + i;
output[LogicalToOffset(logical, ndim, out_shape, out_strides)] =
output_buffer[i];
}
}
processed += current;
}
}

} // namespace

template <typename T>
void GeluUnion(void* workspace, int core_per_cluster, int cluster_count,
cnrtQueue_t queue, const void* input, void* out,
const size_t* input_shape, const ptrdiff_t* input_strides,
const size_t* out_shape, const ptrdiff_t* out_strides,
size_t output_size, int ndim, bool approximate,
bool input_contiguous, bool out_contiguous) {
size_t* device_input_shape = nullptr;
size_t* device_out_shape = nullptr;
ptrdiff_t* device_input_strides = nullptr;
ptrdiff_t* device_out_strides = nullptr;
if (ndim != 0) {
auto* workspace_bytes = static_cast<char*>(workspace);
device_input_shape = reinterpret_cast<size_t*>(workspace_bytes);
device_out_shape = device_input_shape + ndim;
device_input_strides =
reinterpret_cast<ptrdiff_t*>(device_out_shape + ndim);
device_out_strides = device_input_strides + ndim;

CNRT_CHECK(
cnrtMemcpyAsync(device_input_shape, const_cast<size_t*>(input_shape),
ndim * sizeof(size_t), queue, cnrtMemcpyHostToDev));
CNRT_CHECK(cnrtMemcpyAsync(device_out_shape, const_cast<size_t*>(out_shape),
ndim * sizeof(size_t), queue,
cnrtMemcpyHostToDev));
CNRT_CHECK(cnrtMemcpyAsync(
device_input_strides, const_cast<ptrdiff_t*>(input_strides),
ndim * sizeof(ptrdiff_t), queue, cnrtMemcpyHostToDev));
CNRT_CHECK(
cnrtMemcpyAsync(device_out_strides, const_cast<ptrdiff_t*>(out_strides),
ndim * sizeof(ptrdiff_t), queue, cnrtMemcpyHostToDev));
}

cnrtDim3_t kernel_dim;
kernel_dim.x = core_per_cluster;
kernel_dim.y = cluster_count;
kernel_dim.z = 1;
(void)cnrtGetLastError();
GeluKernel<T><<<kernel_dim, cnrtFuncTypeUnion1, queue>>>(
static_cast<const T*>(input), static_cast<T*>(out), device_input_shape,
device_input_strides, device_out_shape, device_out_strides, output_size,
ndim, approximate, input_contiguous, out_contiguous);
CNRT_CHECK(cnrtGetLastError());
}

#define INSTANTIATE_GELU(T) \
template void GeluUnion<T>(void*, int, int, cnrtQueue_t, const void*, void*, \
const size_t*, const ptrdiff_t*, const size_t*, \
const ptrdiff_t*, size_t, int, bool, bool, bool)

INSTANTIATE_GELU(__half);
INSTANTIATE_GELU(__bang_bfloat16);
INSTANTIATE_GELU(float);

#undef INSTANTIATE_GELU

} // namespace infini::ops
2 changes: 2 additions & 0 deletions tests/test_gelu.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,8 @@ def test_gelu(
):
if device == "musa" and dtype == torch.float64:
pytest.skip("MUSA does not support float64 GELU")
if device == "mlu" and dtype == torch.float64:
pytest.skip("Cambricon CNNL does not support float64 GELU")

input = randn_strided(shape, input_strides, dtype=dtype, device=device)
out = (
Expand Down
Loading