diff --git a/src/native/cambricon/ops/relu/kernel.h b/src/native/cambricon/ops/relu/kernel.h new file mode 100644 index 000000000..16dcd2da8 --- /dev/null +++ b/src/native/cambricon/ops/relu/kernel.h @@ -0,0 +1,87 @@ +#ifndef INFINI_OPS_CAMBRICON_RELU_KERNEL_H_ +#define INFINI_OPS_CAMBRICON_RELU_KERNEL_H_ + +#include +#include + +#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 +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 : 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(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>( + {out_type_}, + [&](auto tag) { + using T = typename decltype(tag)::type; + ReluUnion(workspace, core_per_cluster_, cluster_count_, queue, + input.data(), out.data(), input_shape_.data(), + input_strides_.data(), out_strides_.data(), output_size_, + static_cast(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 diff --git a/src/native/cambricon/ops/relu/kernel.mlu b/src/native/cambricon/ops/relu/kernel.mlu new file mode 100644 index 000000000..4e5dd5863 --- /dev/null +++ b/src/native/cambricon/ops/relu/kernel.mlu @@ -0,0 +1,175 @@ +#include +#include + +#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 +__mlu_device__ void ComputeRelu(const T* input, T* output, float* values, + size_t count) { + if constexpr (std::is_same_v) { + __bang_half2float(values, reinterpret_cast(const_cast(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(output), values, count); + } else if constexpr (std::is_same_v) { + __bang_bfloat162float(values, const_cast(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(0) + ? value + : static_cast(0); + } + } +} + +template +__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 +__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(relu_nram_buffer); + auto* output_buffer = input_buffer + block_size; + auto* values = reinterpret_cast(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 +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(workspace); + auto* device_shape = reinterpret_cast(workspace_bytes); + auto* device_input_strides = + reinterpret_cast(device_shape + ndim); + auto* device_out_strides = device_input_strides + ndim; + auto* contiguous_input = reinterpret_cast(device_out_strides + ndim); + + if (ndim != 0) { + CNRT_CHECK(cnrtMemcpyAsync(device_shape, const_cast(shape), + ndim * sizeof(size_t), queue, + cnrtMemcpyHostToDev)); + CNRT_CHECK(cnrtMemcpyAsync( + device_input_strides, const_cast(input_strides), + ndim * sizeof(ptrdiff_t), queue, cnrtMemcpyHostToDev)); + CNRT_CHECK( + cnrtMemcpyAsync(device_out_strides, const_cast(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(input); + if (needs_input_copy) { + if (input_contiguous) { + CNRT_CHECK(cnrtMemcpyAsync(contiguous_input, const_cast(input), + output_size * sizeof(T), queue, + cnrtMemcpyDevToDev)); + } else { + (void)cnrtGetLastError(); + GatherInputKernel<<>>( + 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<<>>( + kernel_input, static_cast(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(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 diff --git a/tests/test_relu.py b/tests/test_relu.py index d3d0c6413..5a2f62d57 100644 --- a/tests/test_relu.py +++ b/tests/test_relu.py @@ -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) @@ -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")],