Skip to content

Commit e5b7ca3

Browse files
committed
Arange is now like the other kernels
1 parent 31985c7 commit e5b7ca3

4 files changed

Lines changed: 39 additions & 37 deletions

File tree

src/tensor/cpu/ops.cpp

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -598,13 +598,15 @@ template Tensor<bfloat16, CPU> cat(const TensorView<bfloat16, CPU>&,
598598
const TensorView<bfloat16, CPU>&, int);
599599
template Tensor<float, CPU> cat(const TensorView<float, CPU>&, const TensorView<float, CPU>&, int);
600600

601+
// elementwise maps
601602
template Tensor<float, CPU> cos(const TensorView<float, CPU>& tensor);
602603
template Tensor<float, CPU> sin(const TensorView<float, CPU>& tensor);
603604
template Tensor<float, CPU> exp(const TensorView<float, CPU>& tensor);
604605

605606
template Tensor<bfloat16, CPU> pow(bfloat16, const TensorView<bfloat16, CPU>&);
606607
template Tensor<float, CPU> pow(const TensorView<float, CPU>&, float);
607608
template Tensor<float, CPU> pow(float, const TensorView<float, CPU>&);
609+
608610
template Tensor<bfloat16, CPU> tril(const TensorView<bfloat16, CPU>&, bool);
609611
template Tensor<int, CPU> tril(const TensorView<int, CPU>&, bool);
610612
template Tensor<bfloat16, CPU> slice(const TensorView<bfloat16, CPU>&, int, size_t, size_t);

src/tensor/cuda/kernels/arange.cu

Lines changed: 27 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
#include "arange.cuh"
2-
#include <cstddef>
2+
#include "utils.cuh"
33

44
namespace tensor::kernels {
55

@@ -14,8 +14,31 @@ __global__ void arange_kernel(DeviceT* out, DeviceT start, DeviceT end, DeviceT
1414
}
1515
}
1616

17-
template __global__ void arange_kernel<Cuda<float>>(Cuda<float>*, Cuda<float>, Cuda<float>, Cuda<float>, size_t);
18-
template __global__ void arange_kernel<Cuda<int>>(Cuda<int>*, Cuda<int>, Cuda<int>, Cuda<int>, size_t);
19-
template __global__ void arange_kernel<Cuda<bfloat16>>(Cuda<bfloat16>*, Cuda<bfloat16>, Cuda<bfloat16>, Cuda<bfloat16>, size_t);
17+
template <typename T> Tensor<T, CUDA> arange(T start, T end, T step) {
18+
auto n_elements = static_cast<size_t>((end - start) / step);
19+
20+
TensorStorage<T, CUDA> storage(n_elements);
21+
22+
Shape shape{n_elements};
23+
24+
Tensor<T, CUDA> out{shape, std::move(storage)};
25+
26+
size_t block_size = cuda::get_block_size(n_elements);
27+
size_t grid_size = cuda::get_grid_size(n_elements, block_size);
28+
29+
// Convert to device-native types for kernel call
30+
auto* device_data = reinterpret_cast<Cuda<T>*>(out.data()); // NOLINT
31+
Cuda<T> device_start = to_device_type(start, CUDA{});
32+
Cuda<T> device_end = to_device_type(end, CUDA{});
33+
Cuda<T> device_step = to_device_type(step, CUDA{});
34+
35+
arange_kernel<Cuda<T>><<<grid_size, block_size>>>(device_data, device_start, device_end, device_step, n_elements);
36+
37+
return out;
38+
}
39+
40+
template Tensor<bfloat16, CUDA> arange(bfloat16 start, bfloat16 end, bfloat16 step);
41+
template Tensor<float, CUDA> arange(float start, float end, float step);
42+
template Tensor<int, CUDA> arange(int start, int end, int step);
2043

2144
} // namespace tensor::kernels

src/tensor/cuda/kernels/arange.cuh

Lines changed: 2 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1,18 +1,12 @@
11
#pragma once
22

3-
#include <cuda_runtime.h>
3+
#include <tensor/tensor.hpp>
44
#include <tensor/device_type.hpp>
5-
#include <cstddef>
65

76
namespace tensor::kernels {
87

98
using namespace dtype;
109

11-
template<typename DeviceT>
12-
__global__ void arange_kernel(DeviceT* out, DeviceT start, DeviceT end, DeviceT step, size_t n);
13-
14-
extern template __global__ void arange_kernel<Cuda<float>>(Cuda<float>*, Cuda<float>, Cuda<float>, Cuda<float>, size_t);
15-
extern template __global__ void arange_kernel<Cuda<int>>(Cuda<int>*, Cuda<int>, Cuda<int>, Cuda<int>, size_t);
16-
extern template __global__ void arange_kernel<Cuda<bfloat16>>(Cuda<bfloat16>*, Cuda<bfloat16>, Cuda<bfloat16>, Cuda<bfloat16>, size_t);
10+
template <typename T> Tensor<T, CUDA> arange(T start, T end, T step);
1711

1812
} // namespace tensor::kernels

src/tensor/cuda/ops.cu

Lines changed: 8 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -20,32 +20,15 @@ namespace tensor {
2020
using namespace dtype;
2121
using namespace device;
2222

23-
template <typename T, typename D> Tensor<T, D> arange(T start, T end, T step) {
24-
auto n_elements = static_cast<size_t>((end - start) / step);
25-
26-
TensorStorage<T, D> storage(n_elements);
27-
28-
Shape shape{n_elements};
29-
30-
Tensor<T, D> out{shape, std::move(storage)};
31-
32-
size_t block_size = cuda::get_block_size(n_elements);
33-
size_t grid_size = cuda::get_grid_size(n_elements, block_size);
34-
35-
// Convert to device-native types for kernel call
36-
auto* device_data = reinterpret_cast<Cuda<T>*>(out.data()); // NOLINT
37-
Cuda<T> device_start = to_device_type(start, D{});
38-
Cuda<T> device_end = to_device_type(end, D{});
39-
Cuda<T> device_step = to_device_type(step, D{});
40-
41-
kernels::arange_kernel<<<grid_size, block_size>>>(device_data, device_start, device_end, device_step, n_elements);
42-
43-
return out;
23+
template <> Tensor<bfloat16, CUDA> arange(bfloat16 start, bfloat16 end, bfloat16 step) {
24+
return kernels::arange(start, end, step);
25+
}
26+
template <> Tensor<float, CUDA> arange(float start, float end, float step) {
27+
return kernels::arange(start, end, step);
28+
}
29+
template <> Tensor<int, CUDA> arange(int start, int end, int step) {
30+
return kernels::arange(start, end, step);
4431
}
45-
46-
template Tensor<int, CUDA> arange(int start, int end, int step);
47-
template Tensor<float, CUDA> arange(float start, float end, float step);
48-
template Tensor<bfloat16, CUDA> arange(bfloat16 start, bfloat16 end, bfloat16 step);
4932

5033
template <typename T, typename D>
5134
void replace_from_(Tensor<T, D>& out, const TensorView<T, D>& input) {

0 commit comments

Comments
 (0)