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
87 changes: 87 additions & 0 deletions src/native/cambricon/ops/relu/kernel.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
#ifndef INFINI_OPS_CAMBRICON_RELU_KERNEL_H_
#define INFINI_OPS_CAMBRICON_RELU_KERNEL_H_

#include <cassert>
#include <cstddef>

#include "base/relu.h"
#include "native/cambricon/cnrt_utils.h"
#include "native/cambricon/common.h"
#include "native/cambricon/data_type_.h"

namespace infini::ops {

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

template <>
class Operator<Relu, Device::Type::kCambricon> : public Relu {
public:
Operator(const Tensor input, Tensor out)
: Relu{input, out}, element_size_{out.element_size()} {
assert(input_type_ != DataType::kFloat64 &&
"`CambriconRelu` does not support float64 because the Cambricon "
"device compiler does not support float64 comparisons.");
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));
}
}

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

void operator()(const Tensor input, Tensor out) const override {
if (output_size_ == 0) {
return;
}

auto queue = static_cast<cnrtQueue_t>(stream_ ? stream_ : 0);
auto* workspace = workspace_ ? workspace_ : default_workspace_;
const auto available_workspace_size =
workspace_ ? workspace_size_in_bytes_ : workspace_size_in_bytes();
assert(available_workspace_size >= workspace_size_in_bytes() &&
"`CambriconRelu` requires a sufficiently large workspace.");
const bool needs_input_copy = NeedsInputCopy(input, out);

DispatchFunc<Device::Type::kCambricon,
List<DataType::kFloat32, DataType::kFloat16,
DataType::kBFloat16, DataType::kInt64, DataType::kInt32,
DataType::kInt16, DataType::kInt8, DataType::kUInt8>>(
{out_type_},
[&](auto tag) {
using T = typename decltype(tag)::type;
ReluUnion<T>(workspace, core_per_cluster_, cluster_count_, queue,
input.data(), out.data(), input_shape_.data(),
input_strides_.data(), out_strides_.data(), output_size_,
static_cast<int>(ndim_), is_input_contiguous_,
is_out_contiguous_, needs_input_copy);
},
"CambriconRelu::operator() - output dispatch");
}

std::size_t workspace_size_in_bytes() const override {
const auto metadata_size = ndim_ * (sizeof(size_t) + 2 * sizeof(ptrdiff_t));
return metadata_size + output_size_ * element_size_;
}

private:
std::size_t element_size_{0};
void* default_workspace_{nullptr};
int core_per_cluster_{0};
int cluster_count_{0};
};

} // namespace infini::ops

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

#include "kernel.h"
#include "native/cambricon/kernel_utils.h"

namespace infini::ops {
namespace {

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

template <typename T>
__mlu_device__ void ComputeRelu(const T* input, T* output, float* values,
size_t count) {
if constexpr (std::is_same_v<T, __half>) {
__bang_half2float(values, reinterpret_cast<half*>(const_cast<T*>(input)),
count);
for (size_t i = 0; i < count; ++i) {
const float value = values[i];
values[i] = value != value || value > 0.0F ? value : 0.0F;
}
__bang_float2half(reinterpret_cast<half*>(output), values, count);
} else if constexpr (std::is_same_v<T, __bang_bfloat16>) {
__bang_bfloat162float(values, const_cast<T*>(input), count);
for (size_t i = 0; i < count; ++i) {
const float value = values[i];
values[i] = value != value || value > 0.0F ? value : 0.0F;
}
__bang_float2bfloat16(output, values, count);
} else {
for (size_t i = 0; i < count; ++i) {
const T value = input[i];
output[i] = value != value || value > static_cast<T>(0)
? value
: static_cast<T>(0);
}
}
}

template <typename T>
__mlu_global__ void GatherInputKernel(const T* input, T* contiguous_input,
const size_t* shape,
const ptrdiff_t* input_strides,
size_t output_size, int ndim) {
const auto range = cambricon::kernel_utils::GetTaskRange(output_size);
for (size_t logical = range.begin; logical < range.end; ++logical) {
contiguous_input[logical] = input[cambricon::kernel_utils::LogicalToOffset(
logical, ndim, shape, input_strides)];
}
}

template <typename T>
__mlu_global__ void ReluKernel(const T* input, T* output, const size_t* shape,
const ptrdiff_t* input_strides,
const ptrdiff_t* out_strides, size_t output_size,
int ndim, bool input_contiguous,
bool out_contiguous) {
const auto range = cambricon::kernel_utils::GetTaskRange(output_size);
if (range.begin >= range.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*>(relu_nram_buffer);
auto* output_buffer = input_buffer + block_size;
auto* values = reinterpret_cast<float*>(output_buffer + block_size);

size_t processed = range.begin;
while (processed < range.end) {
const size_t current =
block_size < range.end - processed ? block_size : range.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[cambricon::kernel_utils::LogicalToOffset(
logical, ndim, shape, input_strides)];
}
}

ComputeRelu(input_buffer, output_buffer, values, current);
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[cambricon::kernel_utils::LogicalToOffset(
logical, ndim, shape, out_strides)] = output_buffer[i];
}
}
processed += current;
}
}

} // namespace

template <typename T>
void ReluUnion(void* workspace, int core_per_cluster, int cluster_count,
cnrtQueue_t queue, const void* input, void* out,
const size_t* shape, const ptrdiff_t* input_strides,
const ptrdiff_t* out_strides, size_t output_size, int ndim,
bool input_contiguous, bool out_contiguous,
bool needs_input_copy) {
auto* workspace_bytes = static_cast<char*>(workspace);
auto* device_shape = reinterpret_cast<size_t*>(workspace_bytes);
auto* device_input_strides =
reinterpret_cast<ptrdiff_t*>(device_shape + ndim);
auto* device_out_strides = device_input_strides + ndim;
auto* contiguous_input = reinterpret_cast<T*>(device_out_strides + ndim);

if (ndim != 0) {
CNRT_CHECK(cnrtMemcpyAsync(device_shape, const_cast<size_t*>(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;

const T* kernel_input = static_cast<const T*>(input);
if (needs_input_copy) {
if (input_contiguous) {
CNRT_CHECK(cnrtMemcpyAsync(contiguous_input, const_cast<void*>(input),
output_size * sizeof(T), queue,
cnrtMemcpyDevToDev));
} else {
(void)cnrtGetLastError();
GatherInputKernel<T><<<kernel_dim, cnrtFuncTypeUnion1, queue>>>(
kernel_input, contiguous_input, device_shape, device_input_strides,
output_size, ndim);
CNRT_CHECK(cnrtGetLastError());
}
kernel_input = contiguous_input;
input_contiguous = true;
}

(void)cnrtGetLastError();
ReluKernel<T><<<kernel_dim, cnrtFuncTypeUnion1, queue>>>(
kernel_input, static_cast<T*>(out), device_shape, device_input_strides,
device_out_strides, output_size, ndim, input_contiguous, out_contiguous);
CNRT_CHECK(cnrtGetLastError());
}

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

INSTANTIATE_RELU(float);
INSTANTIATE_RELU(__half);
INSTANTIATE_RELU(__bang_bfloat16);
INSTANTIATE_RELU(int64_t);
INSTANTIATE_RELU(int32_t);
INSTANTIATE_RELU(int16_t);
INSTANTIATE_RELU(int8_t);
INSTANTIATE_RELU(uint8_t);

#undef INSTANTIATE_RELU

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

input = rand_strided(shape, input_strides, dtype=dtype, device=device)
input.mul_(2).sub_(1)
Expand Down Expand Up @@ -130,6 +132,8 @@ def test_relu_matches_special_value_semantics(
):
if device == "musa" and dtype == torch.float64:
pytest.skip("MUSA does not support float64 ReLU")
if device == "mlu" and dtype == torch.float64:
pytest.skip("Cambricon device code does not support float64 comparisons")

input = torch.tensor(
[float("-inf"), -1.0, -0.0, 0.0, 1.0, float("inf"), float("nan")],
Expand Down
Loading