From 657a1757973ea9e47b77137ea92d2570c15eb676 Mon Sep 17 00:00:00 2001 From: Jacob Domagala Date: Thu, 13 Nov 2025 12:00:00 +0100 Subject: [PATCH 01/13] [#81]: Add fp16 dtype helpers --- python/bindings.cpp | 1 + src/dtype_utils.hpp | 103 ++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 104 insertions(+) create mode 100644 src/dtype_utils.hpp diff --git a/python/bindings.cpp b/python/bindings.cpp index 7914f8e..598f563 100644 --- a/python/bindings.cpp +++ b/python/bindings.cpp @@ -237,6 +237,7 @@ PYBIND11_MODULE(smollnet, m) { }); py::enum_(m, "DataType") + .value("f16", smollnet::DataType::f16) .value("f32", smollnet::DataType::f32) .export_values(); diff --git a/src/dtype_utils.hpp b/src/dtype_utils.hpp new file mode 100644 index 0000000..ba86c54 --- /dev/null +++ b/src/dtype_utils.hpp @@ -0,0 +1,103 @@ +#pragma once + +#include "helpers.hpp" +#include "types.hpp" + +#include + +#include + +namespace smollnet { + +#ifdef __CUDACC__ +#define SMOLLNET_HOST_DEVICE __host__ __device__ +#else +#define SMOLLNET_HOST_DEVICE +#endif + +inline constexpr bool is_supported_float_dtype(DataType dtype) noexcept { + return dtype == DataType::f16 || dtype == DataType::f32; +} + +inline DataType accumulation_dtype(DataType dtype) { + ASSERT(is_supported_float_dtype(dtype), + fmt::format("Unsupported dtype {}", get_name(dtype))); + return DataType::f32; +} + +template struct ScalarTraits; + +template <> struct ScalarTraits { + static constexpr DataType dtype = DataType::f32; + + SMOLLNET_HOST_DEVICE static float to_float(float value) { return value; } + SMOLLNET_HOST_DEVICE static float from_float(float value) { return value; } +}; + +template <> struct ScalarTraits<__half> { + static constexpr DataType dtype = DataType::f16; + + SMOLLNET_HOST_DEVICE static float to_float(__half value) { + return __half2float(value); + } + + SMOLLNET_HOST_DEVICE static __half from_float(float value) { + return __float2half(value); + } +}; + +template +SMOLLNET_HOST_DEVICE float scalar_to_float(T value) { + return ScalarTraits::to_float(value); +} + +template +SMOLLNET_HOST_DEVICE T scalar_from_float(float value) { + return ScalarTraits::from_float(value); +} + +inline float load_scalar(const void *data, DataType dtype, size_t index) { + switch (dtype) { + case DataType::f16: + return scalar_to_float(static_cast(data)[index]); + case DataType::f32: + return static_cast(data)[index]; + default: + ASSERT(false, fmt::format("Unsupported dtype {}", get_name(dtype))); + } + + __builtin_unreachable(); +} + +inline void store_scalar(void *data, DataType dtype, size_t index, + float value) { + switch (dtype) { + case DataType::f16: + static_cast<__half *>(data)[index] = scalar_from_float<__half>(value); + return; + case DataType::f32: + static_cast(data)[index] = value; + return; + default: + ASSERT(false, fmt::format("Unsupported dtype {}", get_name(dtype))); + } + + __builtin_unreachable(); +} + +template void dispatch_float_dtype(DataType dtype, Fn &&fn) { + switch (dtype) { + case DataType::f16: + fn.template operator()<__half>(); + return; + case DataType::f32: + fn.template operator()(); + return; + default: + ASSERT(false, fmt::format("Unsupported dtype {}", get_name(dtype))); + } + + __builtin_unreachable(); +} + +} // namespace smollnet From 552120b9809a27529c303267ffb766fe0cea08a2 Mon Sep 17 00:00:00 2001 From: Jacob Domagala Date: Thu, 13 Nov 2025 12:01:00 +0100 Subject: [PATCH 02/13] [#81]: Add dtype-aware kernel launch APIs --- src/kernels.cuh | 95 ++++++++++++++++++++++++++++--------------------- 1 file changed, 55 insertions(+), 40 deletions(-) diff --git a/src/kernels.cuh b/src/kernels.cuh index b19b3a6..57e1618 100644 --- a/src/kernels.cuh +++ b/src/kernels.cuh @@ -12,7 +12,6 @@ constexpr int32_t ROW_MAJOR = 0; constexpr int32_t COL_MAJOR = 1; constexpr int32_t DEPTH_MAJOR = 2; - struct StrideAndSize { int64_t stride[kMaxTensorDims] = {}; @@ -34,69 +33,85 @@ struct SizeInfo { int64_t b_size[kMaxTensorDims] = {}; }; -enum class WelfordType : uint8_t{ +enum class WelfordType : uint8_t { Mean, PopulationVariance, SampleVariance }; - -void launch_fill(float *ptr, size_t numElems, float val); +void launch_fill(void *ptr, DataType dtype, size_t numElems, float val); void launch_random_init(unsigned long long seed); -void launch_random_fill(void* out, size_t total); -void launch_negative(void *ptr, size_t total); +void launch_random_fill(void *out, DataType dtype, size_t total); +void launch_negative(void *ptr, DataType dtype, size_t total); // Binary OPS -void launch_add(float *out, float *a, float *b, size_t numElems); -void launch_add_strided(void *dst, void *a, void *b, const StrideInfo &s, - size_t total); -void launch_sub(float *out, float *a, float *b, size_t numElems); -void launch_sub_strided(void *out, void *a, void *b, const StrideInfo &s, - size_t total); -void launch_mul(float *out, float *a, float *b, size_t numElems); -void launch_mul_strided(void *dst, void *a, void *b, const StrideInfo &s, - size_t total); -void launch_div(float *out, float *a, float *b, size_t numElems); -void launch_div_strided(void *dst, void *a, void *b, const StrideInfo &s, - size_t total); - -void launch_sum_dim(void *out, void *in, const StrideAndSize &s_input, +void launch_add(void *out, const void *a, const void *b, DataType dtype, + size_t numElems); +void launch_add_strided(void *dst, const void *a, const void *b, + DataType dtype, const StrideInfo &s, size_t total); +void launch_sub(void *out, const void *a, const void *b, DataType dtype, + size_t numElems); +void launch_sub_strided(void *out, const void *a, const void *b, + DataType dtype, const StrideInfo &s, size_t total); +void launch_mul(void *out, const void *a, const void *b, DataType dtype, + size_t numElems); +void launch_mul_strided(void *dst, const void *a, const void *b, + DataType dtype, const StrideInfo &s, size_t total); +void launch_div(void *out, const void *a, const void *b, DataType dtype, + size_t numElems); +void launch_div_strided(void *dst, const void *a, const void *b, + DataType dtype, const StrideInfo &s, size_t total); + +void launch_sum_dim(void *out, const void *in, DataType input_dtype, + const StrideAndSize &s_input, const StrideAndSize &s_output, int64_t dim); -void launch_matmul(void *out, void *left, void *right, +void launch_matmul(void *out, const void *left, const void *right, DataType dtype, const StrideInfo &strides, const SizeInfo &sizes, size_t total); // ACTIVATIONS -void launch_relu(void *out, void *in, size_t total); -void launch_relu_grad(void *out, void *grad_out, void *in, size_t total); +void launch_relu(void *out, const void *in, DataType dtype, size_t total); +void launch_relu_grad(void *out, const void *grad_out, const void *in, + DataType dtype, size_t total); -void launch_gelu(void *out, void *in, size_t total); -void launch_gelu_grad(void *out, void *grad_out, void *in, size_t total); +void launch_gelu(void *out, const void *in, DataType dtype, size_t total); +void launch_gelu_grad(void *out, const void *grad_out, const void *in, + DataType dtype, size_t total); -void launch_tanh(void *out, void *in, size_t total); -void launch_tanh_grad(void *out, void *grad_out, void *in, size_t total); +void launch_tanh(void *out, const void *in, DataType dtype, size_t total); +void launch_tanh_grad(void *out, const void *grad_out, const void *in, + DataType dtype, size_t total); -void launch_sigmoid(void *out, void *in, size_t total); -void launch_sigmoid_grad(void *out, void *grad_out, void *in, size_t total); +void launch_sigmoid(void *out, const void *in, DataType dtype, size_t total); +void launch_sigmoid_grad(void *out, const void *grad_out, const void *in, + DataType dtype, size_t total); -void launch_mse(void *out, void *pred, void *target, size_t total); -void launch_sgd_update(void *p, void *g, float lr, size_t total); -void launch_mse_grad(void *grad, void *pred, void *target, float coeff, - size_t total); +void launch_mse(void *out, DataType out_dtype, const void *pred, + const void *target, DataType input_dtype, size_t total); +void launch_sgd_update(void *p, const void *g, DataType param_dtype, + DataType grad_dtype, float lr, size_t total); +void launch_mse_grad(void *grad, const void *pred, const void *target, + DataType dtype, float coeff, size_t total); // NORM -void launch_mean_2d(void *out, void *in, size_t d0, size_t d1); +void launch_mean_2d(void *out, DataType out_dtype, const void *in, + DataType in_dtype, size_t d0, size_t d1); -void launch_layer_norm(void *out, void *features, void *mean, void *variance, - void *gamma, void *beta, size_t batch_size, +void launch_layer_norm(void *out, const void *features, const void *mean, + const void *variance, const void *gamma, + const void *beta, DataType data_dtype, + DataType stats_dtype, size_t batch_size, size_t num_features); -void launch_layer_norm_grad(void *out, void *normalized_input, - void *scaled_gradient, void *variance, - void *summed_scale, void *summed_scaled_input, +void launch_layer_norm_grad(void *out, const void *normalized_input, + const void *scaled_gradient, const void *variance, + const void *summed_scale, + const void *summed_scaled_input, + DataType data_dtype, DataType stats_dtype, size_t batch_size, size_t num_features); -void launch_welford(void *in, void *out, size_t dim1_len, size_t dim0_len, +void launch_welford(const void *in, DataType in_dtype, void *out, + DataType out_dtype, size_t dim1_len, size_t dim0_len, int32_t dim, WelfordType type); } // namespace smollnet From 43efffa5d779858c051a70f7b257c00904a4fd9a Mon Sep 17 00:00:00 2001 From: Jacob Domagala Date: Thu, 13 Nov 2025 12:02:00 +0100 Subject: [PATCH 03/13] [#81]: Generalize tensor dtype handling --- src/tensor.cpp | 208 +++++++++++++++++++++++++++++++++++++++---------- 1 file changed, 168 insertions(+), 40 deletions(-) diff --git a/src/tensor.cpp b/src/tensor.cpp index b4c34a2..de71287 100644 --- a/src/tensor.cpp +++ b/src/tensor.cpp @@ -1,10 +1,12 @@ #include "tensor.hpp" #include "autograd.hpp" +#include "dtype_utils.hpp" #include "helpers.hpp" #include "kernels.cuh" #include #include +#include #include #include #include @@ -118,9 +120,10 @@ TensorShape broadcast_shape(const Tensor &lhs, const Tensor &rhs, return out_shape; } -using ContiguousBinaryLaunch = void (*)(float *, float *, float *, size_t); -using StridedBinaryLaunch = void (*)(void *, void *, void *, const StrideInfo &, - size_t); +using ContiguousBinaryLaunch = void (*)(void *, const void *, const void *, + DataType, size_t); +using StridedBinaryLaunch = void (*)(void *, const void *, const void *, + DataType, const StrideInfo &, size_t); using CpuBinaryOp = float (*)(float, float); float add_values(float lhs, float rhs) { return lhs + rhs; } @@ -173,25 +176,24 @@ Tensor binary_tensor_op(const Tensor &lhs, const Tensor &rhs, lhs.requires_grad() || rhs.requires_grad()); if (lhs.device() == Device::CPU) { - const auto *lhs_data = static_cast(lhs_view.data()); - const auto *rhs_data = static_cast(rhs_view.data()); - auto *out_data = static_cast(out.data()); - for (size_t idx = 0; idx < out.numel(); ++idx) { int64_t lhs_offset = 0; int64_t rhs_offset = 0; compute_binary_offsets(idx, out_shape, lhs_view.strides(), rhs_view.strides(), out_rank, lhs_offset, rhs_offset); - out_data[idx] = cpu_op(lhs_data[lhs_offset], rhs_data[rhs_offset]); + const float lhs_value = + load_scalar(lhs_view.data(), lhs.dtype(), lhs_offset); + const float rhs_value = + load_scalar(rhs_view.data(), rhs.dtype(), rhs_offset); + store_scalar(out.data(), out.dtype(), idx, cpu_op(lhs_value, rhs_value)); } return out; } if (is_dense_contiguous(lhs_view) && is_dense_contiguous(rhs_view)) { - launch_contiguous(static_cast(out.data()), - static_cast(lhs_view.data()), - static_cast(rhs_view.data()), out.numel()); + launch_contiguous(out.data(), lhs_view.data(), rhs_view.data(), lhs.dtype(), + out.numel()); return out; } @@ -203,17 +205,19 @@ Tensor binary_tensor_op(const Tensor &lhs, const Tensor &rhs, stride_info.b_stride[dim] = rhs_view.strides()[dim]; } - launch_strided(out.data(), lhs_view.data(), rhs_view.data(), stride_info, - out.numel()); + launch_strided(out.data(), lhs_view.data(), rhs_view.data(), lhs.dtype(), + stride_info, out.numel()); return out; } -void append_tensor_values(fmt::memory_buffer &out, const float *data, +void append_tensor_values(fmt::memory_buffer &out, const void *data, + DataType dtype, const TensorShape &sizes, const TensorShape &strides, int64_t rank, int64_t dim, int64_t offset) { if (rank == 0) { - fmt::format_to(std::back_inserter(out), "{:.4f}", data[offset]); + fmt::format_to(std::back_inserter(out), "{:.4f}", + load_scalar(data, dtype, offset)); return; } @@ -221,9 +225,10 @@ void append_tensor_values(fmt::memory_buffer &out, const float *data, for (int64_t i = 0; i < sizes[dim]; ++i) { const int64_t next_offset = offset + i * strides[dim]; if (dim == rank - 1) { - fmt::format_to(std::back_inserter(out), "{:.4f}", data[next_offset]); + fmt::format_to(std::back_inserter(out), "{:.4f}", + load_scalar(data, dtype, next_offset)); } else { - append_tensor_values(out, data, sizes, strides, rank, dim + 1, + append_tensor_values(out, data, dtype, sizes, strides, rank, dim + 1, next_offset); } @@ -260,9 +265,11 @@ Tensor full_like(const Tensor &t, float value, bool requires_grad) { Tensor out = empty(t.dims().data(), t.ndims(), t.dtype(), t.device(), requires_grad); if (t.device() == Device::CUDA) { - launch_fill(static_cast(out.data()), out.numel(), value); + launch_fill(out.data(), out.dtype(), out.numel(), value); } else { - std::fill_n(static_cast(out.data()), out.numel(), value); + for (size_t idx = 0; idx < out.numel(); ++idx) { + store_scalar(out.data(), out.dtype(), idx, value); + } } return out; @@ -328,7 +335,14 @@ void Tensor::zero_grad() const { ASSERT(autograd(), "Tensor doesn't have gradient!"); ASSERT(grad().initialized(), "Gradient is not initialized!"); - launch_fill(static_cast(grad().data()), grad().numel(), 0.0f); + Tensor g = grad(); + if (g.device() == Device::CUDA) { + launch_fill(g.data(), g.dtype(), g.numel(), 0.0f); + } else { + for (size_t idx = 0; idx < g.numel(); ++idx) { + store_scalar(g.data(), g.dtype(), idx, 0.0f); + } + } } bool Tensor::requires_grad() const noexcept { return impl()->requires_grad; } @@ -387,7 +401,7 @@ std::string Tensor::to_string() const { // Could be expensive auto t = cpu(); - const float *raw_data = static_cast(t.data()); + const void *raw_data = t.data(); const auto &sizes = dims(); const auto &stride = strides(); @@ -395,7 +409,7 @@ std::string Tensor::to_string() const { fmt::memory_buffer out; fmt::format_to(std::back_inserter(out), "Tensor: ("); - append_tensor_values(out, raw_data, sizes, stride, ndims(), 0, 0); + append_tensor_values(out, raw_data, dtype(), sizes, stride, ndims(), 0, 0); fmt::format_to(std::back_inserter(out), ")\n"); return fmt::to_string(out); @@ -573,11 +587,35 @@ Tensor matmul(const Tensor &l, const Tensor &r) { ASSERT(l.device() == r.device(), fmt::format("Device mismatch! {} and {}", get_device_name(l.device()), get_device_name(r.device()))); + ASSERT(l.dtype() == r.dtype(), + fmt::format("DType mismatch! {} and {}", get_name(l.dtype()), + get_name(r.dtype()))); bool needs_grad = any_requires_grad({l, r}); Tensor new_tensor = empty({l.dims()[0], r.dims()[1]}, l.dtype(), l.device(), needs_grad); + if (l.device() == Device::CPU) { + for (int64_t row = 0; row < new_tensor.size(0); ++row) { + for (int64_t col = 0; col < new_tensor.size(1); ++col) { + float acc = 0.0f; + for (int64_t k = 0; k < l.size(1); ++k) { + const int64_t lhs_offset = row * l.strides()[0] + k * l.strides()[1]; + const int64_t rhs_offset = k * r.strides()[0] + col * r.strides()[1]; + acc += load_scalar(l.data(), l.dtype(), lhs_offset) * + load_scalar(r.data(), r.dtype(), rhs_offset); + } + + const int64_t out_offset = + row * new_tensor.strides()[0] + col * new_tensor.strides()[1]; + store_scalar(new_tensor.data(), new_tensor.dtype(), out_offset, acc); + } + } + + SetupAutograd(l, r, new_tensor); + return new_tensor; + } + StrideInfo stride_info{}; stride_info.output_size[0] = new_tensor.size(0); stride_info.output_size[1] = new_tensor.size(1); @@ -599,8 +637,8 @@ Tensor matmul(const Tensor &l, const Tensor &r) { size_info.b_size[0] = r.size(0); size_info.b_size[1] = r.size(1); - launch_matmul(new_tensor.data(), l.data(), r.data(), stride_info, size_info, - new_tensor.numel()); + launch_matmul(new_tensor.data(), l.data(), r.data(), l.dtype(), stride_info, + size_info, new_tensor.numel()); SetupAutograd(l, r, new_tensor); @@ -611,7 +649,14 @@ Tensor relu(const Tensor &t) { Tensor new_tensor = empty(t.dims().data(), t.ndims(), t.dtype(), t.device(), t.requires_grad()); - launch_relu(new_tensor.data(), t.data(), t.numel()); + if (t.device() == Device::CUDA) { + launch_relu(new_tensor.data(), t.data(), t.dtype(), t.numel()); + } else { + for (size_t idx = 0; idx < t.numel(); ++idx) { + store_scalar(new_tensor.data(), new_tensor.dtype(), idx, + std::max(load_scalar(t.data(), t.dtype(), idx), 0.0f)); + } + } SetupAutograd(new_tensor, t); @@ -622,7 +667,19 @@ Tensor gelu(const Tensor &t) { Tensor new_tensor = empty(t.dims().data(), t.ndims(), t.dtype(), t.device(), t.requires_grad()); - launch_gelu(new_tensor.data(), t.data(), t.numel()); + if (t.device() == Device::CUDA) { + launch_gelu(new_tensor.data(), t.data(), t.dtype(), t.numel()); + } else { + constexpr float sqrt_2_over_pi = 0.7978845608f; + for (size_t idx = 0; idx < t.numel(); ++idx) { + const float x = load_scalar(t.data(), t.dtype(), idx); + const float value = + 0.5f * x * + (1.0f + + std::tanh(sqrt_2_over_pi * (x + 0.044715f * x * x * x))); + store_scalar(new_tensor.data(), new_tensor.dtype(), idx, value); + } + } SetupAutograd(new_tensor, t); return new_tensor; @@ -632,7 +689,14 @@ Tensor tanh(const Tensor &t) { Tensor new_tensor = empty(t.dims().data(), t.ndims(), t.dtype(), t.device(), t.requires_grad()); - launch_tanh(new_tensor.data(), t.data(), t.numel()); + if (t.device() == Device::CUDA) { + launch_tanh(new_tensor.data(), t.data(), t.dtype(), t.numel()); + } else { + for (size_t idx = 0; idx < t.numel(); ++idx) { + store_scalar(new_tensor.data(), new_tensor.dtype(), idx, + std::tanh(load_scalar(t.data(), t.dtype(), idx))); + } + } SetupAutograd(new_tensor, t); return new_tensor; } @@ -641,7 +705,15 @@ Tensor sigmoid(const Tensor &t) { Tensor new_tensor = empty(t.dims().data(), t.ndims(), t.dtype(), t.device(), t.requires_grad()); - launch_sigmoid(new_tensor.data(), t.data(), t.numel()); + if (t.device() == Device::CUDA) { + launch_sigmoid(new_tensor.data(), t.data(), t.dtype(), t.numel()); + } else { + for (size_t idx = 0; idx < t.numel(); ++idx) { + const float x = load_scalar(t.data(), t.dtype(), idx); + store_scalar(new_tensor.data(), new_tensor.dtype(), idx, + 1.0f / (1.0f + std::exp(-x))); + } + } SetupAutograd(new_tensor, t); return new_tensor; } @@ -649,7 +721,8 @@ Tensor sigmoid(const Tensor &t) { Tensor sum(const Tensor &t, int64_t dim, bool keep_dim) { auto dims = t.dims(); auto new_rank = keep_dim ? t.ndims() : t.ndims() - 1; - auto data_type = t.dtype(); + auto input_type = t.dtype(); + auto output_type = accumulation_dtype(input_type); auto device = t.device(); ASSERT(dim < t.ndims(), @@ -673,7 +746,7 @@ Tensor sum(const Tensor &t, int64_t dim, bool keep_dim) { } Tensor new_tensor = - zeros(out_dims.data(), new_rank, data_type, device, t.requires_grad()); + zeros(out_dims.data(), new_rank, output_type, device, t.requires_grad()); auto *srcp = t.data(); auto *dst = new_tensor.data(); @@ -689,7 +762,37 @@ Tensor sum(const Tensor &t, int64_t dim, bool keep_dim) { copy_shape_to_kernel_array(new_tensor.strides(), s_output.stride, s_output.rank); - launch_sum_dim(dst, srcp, s_input, s_output, dim); + if (device == Device::CPU) { + for (size_t linear = 0; linear < new_tensor.numel(); ++linear) { + size_t remaining = linear; + int64_t input_base_offset = 0; + int64_t output_offset = 0; + + for (int64_t out_dim = s_output.rank - 1; out_dim >= 0; --out_dim) { + const int64_t coord = remaining % s_output.size[out_dim]; + remaining /= s_output.size[out_dim]; + + output_offset += coord * s_output.stride[out_dim]; + + const bool kept_dim = s_output.rank == s_input.rank; + const int64_t input_dim = + kept_dim || out_dim < dim ? out_dim : out_dim + 1; + input_base_offset += coord * s_input.stride[input_dim]; + } + + float acc = 0.0f; + for (int64_t reduce_idx = 0; reduce_idx < s_input.size[dim]; + ++reduce_idx) { + const int64_t input_offset = + input_base_offset + reduce_idx * s_input.stride[dim]; + acc += load_scalar(srcp, input_type, input_offset); + } + + store_scalar(dst, output_type, output_offset, acc); + } + } else { + launch_sum_dim(dst, srcp, input_type, s_input, s_output, dim); + } return new_tensor; } @@ -706,10 +809,31 @@ Tensor div(const Tensor &left, const Tensor &right) { return left.div(right); } Tensor mse(const Tensor &pred, const Tensor &target) { ASSERT(pred.dims() == target.dims(), ""); + ASSERT(pred.device() == target.device(), + fmt::format("Device mismatch! {} and {}", + get_device_name(pred.device()), + get_device_name(target.device()))); + ASSERT(pred.dtype() == target.dtype(), + fmt::format("DType mismatch! {} and {}", get_name(pred.dtype()), + get_name(target.dtype()))); bool requires_grad = any_requires_grad({pred, target}); - auto new_tensor = zeros({1}, pred.dtype(), pred.device(), requires_grad); - launch_mse(new_tensor.data(), pred.data(), target.data(), pred.numel()); + auto new_tensor = + zeros({1}, accumulation_dtype(pred.dtype()), pred.device(), requires_grad); + + if (pred.device() == Device::CPU) { + float acc = 0.0f; + for (size_t idx = 0; idx < pred.numel(); ++idx) { + const float diff = load_scalar(pred.data(), pred.dtype(), idx) - + load_scalar(target.data(), target.dtype(), idx); + acc += diff * diff; + } + store_scalar(new_tensor.data(), new_tensor.dtype(), 0, + acc / static_cast(pred.numel())); + } else { + launch_mse(new_tensor.data(), new_tensor.dtype(), pred.data(), + target.data(), pred.dtype(), pred.numel()); + } SetupAutograd(pred, target, new_tensor); return new_tensor; @@ -786,15 +910,18 @@ Tensor empty(const int64_t *dims, size_t rank, DataType t, Device d, ASSERT(rank <= kMaxTensorDims, fmt::format("Tensor rank {} exceeds max rank {}", rank, kMaxTensorDims)); + ASSERT(is_supported_float_dtype(t), + fmt::format("Unsupported dtype {}. Supported dtypes are f16 and f32", + get_name(t))); auto storage = std::make_shared(); - float *ptr; + void *ptr = nullptr; size_t bytes = element_size(t) * product(dims, rank); if (d == Device::CUDA) { CHECK_CUDA(cudaMalloc(&ptr, bytes)); } else { - ptr = static_cast(malloc(bytes)); + ptr = malloc(bytes); } storage->ptr = ptr; @@ -827,9 +954,11 @@ Tensor ones(const int64_t *dims, size_t rank, DataType t, Device d, auto tensor = empty(dims, rank, t, d, requires_grad); if (d == Device::CUDA) { - launch_fill(static_cast(tensor.data()), tensor.numel(), 1.0f); + launch_fill(tensor.data(), tensor.dtype(), tensor.numel(), 1.0f); } else { - std::fill_n(static_cast(tensor.data()), tensor.numel(), 1.0f); + for (size_t idx = 0; idx < tensor.numel(); ++idx) { + store_scalar(tensor.data(), tensor.dtype(), idx, 1.0f); + } } return Tensor{tensor}; @@ -845,13 +974,12 @@ Tensor rand(const int64_t *dims, size_t rank, DataType t, Device d, auto tensor = empty(dims, rank, t, d, requires_grad); if (d == Device::CUDA) { - launch_random_fill(tensor.data(), tensor.numel()); + launch_random_fill(tensor.data(), tensor.dtype(), tensor.numel()); } else { - auto *data = static_cast(tensor.data()); std::uniform_real_distribution dist(0.0f, 1.0f); auto &generator = cpu_random_generator(); for (size_t i = 0; i < tensor.numel(); ++i) { - data[i] = dist(generator); + store_scalar(tensor.data(), tensor.dtype(), i, dist(generator)); } } From 4c204f814e869a15db34ab365bf6dc97bff25e4d Mon Sep 17 00:00:00 2001 From: Jacob Domagala Date: Thu, 13 Nov 2025 12:03:00 +0100 Subject: [PATCH 04/13] [#81]: Dispatch elementwise kernels by dtype --- src/operators.cu | 255 +++++++++++++++++++++-------------------------- src/prng.cu | 22 ++-- 2 files changed, 125 insertions(+), 152 deletions(-) diff --git a/src/operators.cu b/src/operators.cu index b399c1d..e7f1d4f 100644 --- a/src/operators.cu +++ b/src/operators.cu @@ -1,3 +1,4 @@ +#include "dtype_utils.hpp" #include "helpers.hpp" #include "kernels.cuh" @@ -5,6 +6,24 @@ namespace smollnet { +namespace { + +struct AddOp { + __device__ float operator()(float lhs, float rhs) const { return lhs + rhs; } +}; + +struct SubOp { + __device__ float operator()(float lhs, float rhs) const { return lhs - rhs; } +}; + +struct MulOp { + __device__ float operator()(float lhs, float rhs) const { return lhs * rhs; } +}; + +struct DivOp { + __device__ float operator()(float lhs, float rhs) const { return lhs / rhs; } +}; + __device__ __forceinline__ void compute_strided_offsets(size_t idx, const StrideInfo &s, int64_t &offA, @@ -21,195 +40,145 @@ __device__ __forceinline__ void compute_strided_offsets(size_t idx, } } -__global__ void negative_kernel(float *ptr, size_t total) { - auto idx = threadIdx.x + blockDim.x * blockIdx.x; - if (idx < total) - ptr[idx] *= -1.0f; -} - -void launch_negative(void *ptr, size_t total) { - dim3 block = 256; - dim3 grid = (block.x + total - 1) / block.x; - - negative_kernel<<>>(static_cast(ptr), total); -} - -template __global__ void fill_kernel(T *data, size_t n, T value) { - size_t idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < n) - data[idx] = value; -} - -void launch_fill(float *ptr, size_t numElems, float val) { - dim3 block(256); - dim3 grid((numElems + block.x - 1) / block.x); - fill_kernel<<>>(ptr, numElems, val); - CHECK_CUDA(cudaGetLastError()); +template __global__ void negative_kernel(T *ptr, size_t total) { + const auto idx = threadIdx.x + blockDim.x * blockIdx.x; + if (idx < total) { + ptr[idx] = scalar_from_float(-scalar_to_float(ptr[idx])); + } } -template -__global__ void add_kernel(T *out, T *left, T *right, size_t n) { - size_t idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < n) - out[idx] = left[idx] + right[idx]; +template __global__ void fill_kernel(T *data, size_t n, + float value) { + const size_t idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + data[idx] = scalar_from_float(value); + } } -void launch_add(float *out, float *left, float *right, size_t numElems) { - dim3 block(256); - dim3 grid((numElems + block.x - 1) / block.x); - add_kernel<<>>(out, left, right, numElems); - CHECK_CUDA(cudaGetLastError()); +template +__global__ void binary_kernel(T *out, const T *left, const T *right, size_t n, + Op op) { + const size_t idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + const float lhs = scalar_to_float(left[idx]); + const float rhs = scalar_to_float(right[idx]); + out[idx] = scalar_from_float(op(lhs, rhs)); + } } -__global__ void add_strided_kernel(float *__restrict__ out, - const float *__restrict__ a, - const float *__restrict__ b, StrideInfo s, - size_t total) { - size_t idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx >= total) +template +__global__ void binary_strided_kernel(T *__restrict__ out, + const T *__restrict__ a, + const T *__restrict__ b, StrideInfo s, + size_t total, Op op) { + const size_t idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx >= total) { return; + } int64_t offA = 0; int64_t offB = 0; compute_strided_offsets(idx, s, offA, offB); - out[idx] = a[offA] + b[offB]; + const float lhs = scalar_to_float(a[offA]); + const float rhs = scalar_to_float(b[offB]); + out[idx] = scalar_from_float(op(lhs, rhs)); } -void launch_add_strided(void *dst, void *a, void *b, const StrideInfo &s, - size_t total) { - dim3 blk(256); - dim3 grd((total + blk.x - 1) / blk.x); +template +void launch_binary(void *out, const void *left, const void *right, + DataType dtype, size_t numElems, Op op) { + dim3 block(256); + dim3 grid((numElems + block.x - 1) / block.x); + + dispatch_float_dtype(dtype, [&]() { + binary_kernel<<>>(static_cast(out), + static_cast(left), + static_cast(right), + numElems, op); + }); - add_strided_kernel<<>>(static_cast(dst), - static_cast(a), - static_cast(b), s, total); CHECK_CUDA(cudaGetLastError()); } -template -__global__ void mul_kernel(T *out, T *left, T scalar, size_t n) { - size_t idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < n) - out[idx] = left[idx] * scalar; -} +template +void launch_binary_strided(void *out, const void *left, const void *right, + DataType dtype, const StrideInfo &s, size_t total, + Op op) { + dim3 block(256); + dim3 grid((total + block.x - 1) / block.x); -template -__global__ void mul_kernel(T *out, T *left, T *right, size_t n) { - size_t idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < n) - out[idx] = left[idx] * right[idx]; -} + dispatch_float_dtype(dtype, [&]() { + binary_strided_kernel<<>>( + static_cast(out), static_cast(left), + static_cast(right), s, total, op); + }); -void launch_mul(float *out, float *left, float *right, size_t numElems) { - dim3 block(256); - dim3 grid((numElems + block.x - 1) / block.x); - mul_kernel<<>>(out, left, right, numElems); CHECK_CUDA(cudaGetLastError()); } -__global__ void mul_strided_kernel(float *__restrict__ out, - const float *__restrict__ a, - const float *__restrict__ b, StrideInfo s, - size_t total) { - size_t idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx >= total) - return; - - int64_t offA = 0; - int64_t offB = 0; - compute_strided_offsets(idx, s, offA, offB); +} // namespace - out[idx] = a[offA] * b[offB]; -} +void launch_negative(void *ptr, DataType dtype, size_t total) { + dim3 block = 256; + dim3 grid = (block.x + total - 1) / block.x; -void launch_mul_strided(void *dst, void *a, void *b, const StrideInfo &s, - size_t total) { - dim3 blk(256); - dim3 grd((total + blk.x - 1) / blk.x); + dispatch_float_dtype(dtype, [&]() { + negative_kernel<<>>(static_cast(ptr), total); + }); - mul_strided_kernel<<>>(static_cast(dst), - static_cast(a), - static_cast(b), s, total); CHECK_CUDA(cudaGetLastError()); } -template -__global__ void sub_kernel(T *out, T *left, T *right, size_t n) { - size_t idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < n) - out[idx] = left[idx] - right[idx]; -} - -void launch_sub(float *out, float *a, float *b, size_t numElems) { +void launch_fill(void *ptr, DataType dtype, size_t numElems, float val) { dim3 block(256); dim3 grid((numElems + block.x - 1) / block.x); - sub_kernel<<>>(out, a, b, numElems); - CHECK_CUDA(cudaGetLastError()); -} -__global__ void sub_strided_kernel(float *out, float *a, float *b, StrideInfo s, - size_t total) { - auto idx = threadIdx.x + blockDim.x * blockIdx.x; - - if (idx >= total) - return; - - int64_t offA = 0; - int64_t offB = 0; - compute_strided_offsets(idx, s, offA, offB); + dispatch_float_dtype(dtype, [&]() { + fill_kernel<<>>(static_cast(ptr), numElems, val); + }); - out[idx] = a[offA] - b[offB]; + CHECK_CUDA(cudaGetLastError()); } -void launch_sub_strided(void *out, void *a, void *b, const StrideInfo &s, - size_t total) { - dim3 block = 256; - dim3 grid = (total + block.x - 1) / block.x; - sub_strided_kernel<<>>(static_cast(out), - static_cast(a), - static_cast(b), s, total); - CHECK_CUDA(cudaGetLastError()); +void launch_add(void *out, const void *left, const void *right, DataType dtype, + size_t numElems) { + launch_binary(out, left, right, dtype, numElems, AddOp{}); } -template -__global__ void div_kernel(T *out, T *left, T *right, size_t n) { - size_t idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < n) - out[idx] = left[idx] / right[idx]; +void launch_add_strided(void *dst, const void *a, const void *b, + DataType dtype, const StrideInfo &s, size_t total) { + launch_binary_strided(dst, a, b, dtype, s, total, AddOp{}); } -void launch_div(float *out, float *a, float *b, size_t numElems) { - dim3 block(256); - dim3 grid((numElems + block.x - 1) / block.x); - div_kernel<<>>(out, a, b, numElems); - CHECK_CUDA(cudaGetLastError()); +void launch_mul(void *out, const void *left, const void *right, DataType dtype, + size_t numElems) { + launch_binary(out, left, right, dtype, numElems, MulOp{}); } -__global__ void div_strided_kernel(float *__restrict__ out, - const float *__restrict__ a, - const float *__restrict__ b, StrideInfo s, - size_t total) { - size_t idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx >= total) - return; +void launch_mul_strided(void *dst, const void *a, const void *b, + DataType dtype, const StrideInfo &s, size_t total) { + launch_binary_strided(dst, a, b, dtype, s, total, MulOp{}); +} - int64_t offA = 0; - int64_t offB = 0; - compute_strided_offsets(idx, s, offA, offB); +void launch_sub(void *out, const void *a, const void *b, DataType dtype, + size_t numElems) { + launch_binary(out, a, b, dtype, numElems, SubOp{}); +} - out[idx] = a[offA] / b[offB]; +void launch_sub_strided(void *out, const void *a, const void *b, + DataType dtype, const StrideInfo &s, size_t total) { + launch_binary_strided(out, a, b, dtype, s, total, SubOp{}); } -void launch_div_strided(void *dst, void *a, void *b, const StrideInfo &s, - size_t total) { - dim3 blk(256); - dim3 grd((total + blk.x - 1) / blk.x); +void launch_div(void *out, const void *a, const void *b, DataType dtype, + size_t numElems) { + launch_binary(out, a, b, dtype, numElems, DivOp{}); +} - div_strided_kernel<<>>(static_cast(dst), - static_cast(a), - static_cast(b), s, total); - CHECK_CUDA(cudaGetLastError()); +void launch_div_strided(void *dst, const void *a, const void *b, + DataType dtype, const StrideInfo &s, size_t total) { + launch_binary_strided(dst, a, b, dtype, s, total, DivOp{}); } } // namespace smollnet diff --git a/src/prng.cu b/src/prng.cu index f3a4023..73e0e5f 100644 --- a/src/prng.cu +++ b/src/prng.cu @@ -1,3 +1,4 @@ +#include "dtype_utils.hpp" #include "helpers.hpp" #include @@ -15,7 +16,8 @@ std::size_t g_num_states = 0; } // namespace -__global__ void random_fill_kernel(float *out, +template +__global__ void random_fill_kernel(T *out, curandStatePhilox4_32_10_t *states, std::size_t total, std::size_t num_states) { @@ -34,19 +36,19 @@ __global__ void random_fill_kernel(float *out, while (idx < total) { const float4 r = curand_uniform4(&local_state); - out[idx] = r.x; + out[idx] = scalar_from_float(r.x); idx += stride; if (idx >= total) break; - out[idx] = r.y; + out[idx] = scalar_from_float(r.y); idx += stride; if (idx >= total) break; - out[idx] = r.z; + out[idx] = scalar_from_float(r.z); idx += stride; if (idx >= total) break; - out[idx] = r.w; + out[idx] = scalar_from_float(r.w); idx += stride; } @@ -98,7 +100,7 @@ void launch_random_init(unsigned long long seed) { CHECK_CUDA(cudaGetLastError()); } -void launch_random_fill(void *out, std::size_t total) { +void launch_random_fill(void *out, DataType dtype, std::size_t total) { if (d_states == nullptr || g_num_states == 0) { launch_random_init(1234ULL); } @@ -106,10 +108,12 @@ void launch_random_fill(void *out, std::size_t total) { dim3 block_size(kBlockSize); dim3 grid_size((g_num_states + block_size.x - 1) / block_size.x); - random_fill_kernel<<>>(static_cast(out), - d_states, total, g_num_states); + dispatch_float_dtype(dtype, [&]() { + random_fill_kernel<<>>( + static_cast(out), d_states, total, g_num_states); + }); CHECK_CUDA(cudaGetLastError()); } -} // namespace smollnet \ No newline at end of file +} // namespace smollnet From 309f6ff87eedea8e0185da49e82dee4346a901a8 Mon Sep 17 00:00:00 2001 From: Jacob Domagala Date: Thu, 13 Nov 2025 12:04:00 +0100 Subject: [PATCH 05/13] [#81]: Accumulate tensor sums in fp32 --- src/sum.cu | 170 ++++++++++++++++++++++++++++++++--------------------- 1 file changed, 103 insertions(+), 67 deletions(-) diff --git a/src/sum.cu b/src/sum.cu index 78790df..59ea6e1 100644 --- a/src/sum.cu +++ b/src/sum.cu @@ -1,6 +1,9 @@ +#include "dtype_utils.hpp" #include "helpers.hpp" #include "kernels.cuh" +#include + namespace smollnet { namespace { @@ -20,11 +23,9 @@ StrideAndSize padded_rank3(StrideAndSize shape) { return shape; } -} // namespace - -template +template __global__ void -warp_level_sum(const float *__restrict__ in, float *__restrict__ out, +warp_level_sum(const InT *__restrict__ in, float *__restrict__ out, size_t dim_size, int64_t s0_in, int64_t s1_in, int64_t s2_in, int64_t s0_out, int64_t s1_out, int64_t s2_out, size_t n) { // We always launch this kernel as 1D block @@ -61,7 +62,7 @@ warp_level_sum(const float *__restrict__ in, float *__restrict__ out, in_bounds = (idx < n and depth < dim_size); } - float v = in_bounds ? in[idx] : 0.0f; + float v = in_bounds ? scalar_to_float(in[idx]) : 0.0f; __shared__ float sMem[BLOCK_DIM / 32]; @@ -86,10 +87,11 @@ warp_level_sum(const float *__restrict__ in, float *__restrict__ out, } } -template -__global__ void strided_sum_2d(const float *__restrict__ in, +template +__global__ void strided_sum_2d(const InT *__restrict__ in, float *__restrict__ out, int64_t dim_len, - int64_t s0, int64_t s1, int64_t outer_dim_len) { + int64_t s0, int64_t s1, + int64_t outer_dim_len) { int64_t col; int64_t row; @@ -126,15 +128,16 @@ __global__ void strided_sum_2d(const float *__restrict__ in, #pragma unroll for (int elem = 0; elem < VEC_LEN; elem++) { const int64_t my_idx = base_idx + elem * elem_stride; - const float my_val = (base_vec_idx + elem) < dim_len ? in[my_idx] : 0.0f; + const float my_val = + (base_vec_idx + elem) < dim_len ? scalar_to_float(in[my_idx]) : 0.0f; acc += my_val; } atomicAdd(out + outer_idx, acc); } -template -__global__ void strided_sum_3d(const float *__restrict__ in, +template +__global__ void strided_sum_3d(const InT *__restrict__ in, float *__restrict__ out, int64_t dim_len, int64_t s0, int64_t s1, int64_t s2, int64_t main_axis_max_len) { @@ -184,15 +187,17 @@ __global__ void strided_sum_3d(const float *__restrict__ in, #pragma unroll for (int elem = 0; elem < VEC_LEN; elem++) { const int64_t my_idx = base_idx + elem * stride; - const float my_val = - (seconday_axis * VEC_LEN + elem) < dim_len ? in[my_idx] : 0.0f; + const float my_val = (seconday_axis * VEC_LEN + elem) < dim_len + ? scalar_to_float(in[my_idx]) + : 0.0f; acc += my_val; } atomicAdd(out + out_idx, acc); } -void launch_sum_dim0(void *out, void *in, const StrideAndSize &s_input, +template +void launch_sum_dim0(float *out, const void *in, const StrideAndSize &s_input, const StrideAndSize &s_output) { const auto d0 = s_input.size[0]; @@ -201,6 +206,7 @@ void launch_sum_dim0(void *out, void *in, const StrideAndSize &s_input, const int64_t dim_len = d0; constexpr size_t BLOCK = 256; + const auto *typed_in = static_cast(in); // Contigious memory access -> warp level reduce! if (s_input.stride[0] == 1) { @@ -209,25 +215,25 @@ void launch_sum_dim0(void *out, void *in, const StrideAndSize &s_input, if (s_input.rank == 1) { dim3 grid((BLOCK + d0 - 1) / BLOCK, 1, 1); - warp_level_sum<<>>( - static_cast(in), static_cast(out), dim_len, - s_input.stride[0], s_input.stride[0], s_input.stride[0], - s_output.stride[0], s_output.stride[0], s_output.stride[0], total); + warp_level_sum<<>>( + typed_in, out, dim_len, s_input.stride[0], s_input.stride[0], + s_input.stride[0], s_output.stride[0], s_output.stride[0], + s_output.stride[0], total); } else if (s_input.rank == 2) { dim3 grid((BLOCK + d0 - 1) / BLOCK, d1, 1); - warp_level_sum<<>>( - static_cast(in), static_cast(out), dim_len, - s_input.stride[0], s_input.stride[0], s_input.stride[1], - s_output.stride[0], s_output.stride[0], s_output.stride[1], total); + warp_level_sum<<>>( + typed_in, out, dim_len, s_input.stride[0], s_input.stride[0], + s_input.stride[1], s_output.stride[0], s_output.stride[0], + s_output.stride[1], total); } else { dim3 grid((BLOCK + d0 - 1) / BLOCK, d1, d2); - warp_level_sum<<>>( - static_cast(in), static_cast(out), dim_len, - s_input.stride[0], s_input.stride[1], s_input.stride[2], - s_output.stride[0], s_output.stride[1], s_output.stride[2], total); + warp_level_sum<<>>( + typed_in, out, dim_len, s_input.stride[0], s_input.stride[1], + s_input.stride[2], s_output.stride[0], s_output.stride[1], + s_output.stride[2], total); } } else { constexpr int32_t VEC_LEN = 2; @@ -236,22 +242,22 @@ void launch_sum_dim0(void *out, void *in, const StrideAndSize &s_input, if (s_input.rank == 2) { grid = dim3((BLOCK + d1 - 1) / BLOCK, (VEC_LEN + d0 - 1) / VEC_LEN, 1); - strided_sum_2d<<>>( - static_cast(in), static_cast(out), dim_len, - s_input.stride[0], s_input.stride[1], d1); + strided_sum_2d<<>>( + typed_in, out, dim_len, s_input.stride[0], s_input.stride[1], d1); } else { grid = dim3((BLOCK + d2 - 1) / BLOCK, d1, (VEC_LEN + d0 - 1) / VEC_LEN); - strided_sum_3d<<>>( - static_cast(in), static_cast(out), dim_len, - s_input.stride[0], s_input.stride[1], s_input.stride[2], d2); + strided_sum_3d<<>>( + typed_in, out, dim_len, s_input.stride[0], s_input.stride[1], + s_input.stride[2], d2); } } CHECK_CUDA(cudaGetLastError()); } -void launch_sum_dim1(void *out, void *in, const StrideAndSize &s_input, +template +void launch_sum_dim1(float *out, const void *in, const StrideAndSize &s_input, const StrideAndSize &s_output) { const auto d0 = s_input.size[0]; const auto d1 = s_input.size[1]; @@ -259,6 +265,7 @@ void launch_sum_dim1(void *out, void *in, const StrideAndSize &s_input, const int64_t dim_len = d1; constexpr size_t BLOCK = 256; + const auto *typed_in = static_cast(in); // Contigious memory access -> warp level reduce! if (s_input.stride[1] == 1) { @@ -267,18 +274,18 @@ void launch_sum_dim1(void *out, void *in, const StrideAndSize &s_input, if (s_input.rank == 2) { dim3 grid((BLOCK + d1 - 1) / BLOCK, d0, 1); - warp_level_sum<<>>( - static_cast(in), static_cast(out), dim_len, - s_input.stride[0], s_input.stride[0], s_input.stride[1], - s_output.stride[0], s_output.stride[0], s_output.stride[1], total); + warp_level_sum<<>>( + typed_in, out, dim_len, s_input.stride[0], s_input.stride[0], + s_input.stride[1], s_output.stride[0], s_output.stride[0], + s_output.stride[1], total); } else { // Transposed dim3 grid((BLOCK + d1 - 1) / BLOCK, d2, d0); - warp_level_sum<<>>( - static_cast(in), static_cast(out), dim_len, - s_input.stride[0], s_input.stride[1], s_input.stride[2], - s_output.stride[0], s_output.stride[1], s_output.stride[2], total); + warp_level_sum<<>>( + typed_in, out, dim_len, s_input.stride[0], s_input.stride[1], + s_input.stride[2], s_output.stride[0], s_output.stride[1], + s_output.stride[2], total); } } else { @@ -287,22 +294,22 @@ void launch_sum_dim1(void *out, void *in, const StrideAndSize &s_input, if (s_input.rank == 2) { dim3 grid((BLOCK + d0 - 1) / BLOCK, (VEC_LEN + d1 - 1) / VEC_LEN, 1); - strided_sum_2d<<>>( - static_cast(in), static_cast(out), dim_len, - s_input.stride[0], s_input.stride[1], d0); + strided_sum_2d<<>>( + typed_in, out, dim_len, s_input.stride[0], s_input.stride[1], d0); } else { dim3 grid((BLOCK + d2 - 1) / BLOCK, (VEC_LEN + d1 - 1) / VEC_LEN, d0); - strided_sum_3d<<>>( - static_cast(in), static_cast(out), dim_len, - s_input.stride[0], s_input.stride[1], s_input.stride[2], d2); + strided_sum_3d<<>>( + typed_in, out, dim_len, s_input.stride[0], s_input.stride[1], + s_input.stride[2], d2); } } CHECK_CUDA(cudaGetLastError()); } -void launch_sum_dim2(void *out, void *in, const StrideAndSize &s_input, +template +void launch_sum_dim2(float *out, const void *in, const StrideAndSize &s_input, const StrideAndSize &s_output) { const auto d0 = s_input.size[0]; @@ -310,29 +317,31 @@ void launch_sum_dim2(void *out, void *in, const StrideAndSize &s_input, const auto d2 = s_input.size[2]; constexpr size_t BLOCK = 256; + const auto *typed_in = static_cast(in); if (s_input.stride[2] == 1) { const auto total = d0 * d1 * d2; dim3 grid((BLOCK + d2 - 1) / BLOCK, d1, d0); - warp_level_sum<<>>( - static_cast(in), static_cast(out), d2, - s_input.stride[0], s_input.stride[1], s_input.stride[2], - s_output.stride[0], s_output.stride[1], s_output.stride[2], total); + warp_level_sum<<>>( + typed_in, out, d2, s_input.stride[0], s_input.stride[1], + s_input.stride[2], s_output.stride[0], s_output.stride[1], + s_output.stride[2], total); } else { constexpr int32_t VEC_LEN = 64; dim3 grid((d2 + BLOCK - 1) / BLOCK, (d1 + VEC_LEN - 1) / VEC_LEN, d0); - strided_sum_3d<<>>( - static_cast(in), static_cast(out), d2, - s_input.stride[0], s_input.stride[1], s_input.stride[2], d1); + strided_sum_3d<<>>( + typed_in, out, d2, s_input.stride[0], s_input.stride[1], + s_input.stride[2], d1); } CHECK_CUDA(cudaGetLastError()); } -__global__ void generic_sum_dim_kernel(const float *__restrict__ in, +template +__global__ void generic_sum_dim_kernel(const InT *__restrict__ in, float *__restrict__ out, const StrideAndSize s_input, const StrideAndSize s_output, @@ -365,10 +374,12 @@ __global__ void generic_sum_dim_kernel(const float *__restrict__ in, output_offset += coord * s_output.stride[output_dim]; } - atomicAdd(out + output_offset, in[input_offset]); + atomicAdd(out + output_offset, scalar_to_float(in[input_offset])); } -void launch_generic_sum_dim(void *out, void *in, const StrideAndSize &s_input, +template +void launch_generic_sum_dim(float *out, const void *in, + const StrideAndSize &s_input, const StrideAndSize &s_output, int64_t dim) { const size_t total = shape_numel(s_input.size, s_input.rank); if (total == 0) { @@ -378,33 +389,58 @@ void launch_generic_sum_dim(void *out, void *in, const StrideAndSize &s_input, constexpr int block = 256; const int grid = static_cast((total + block - 1) / block); - generic_sum_dim_kernel<<>>( - static_cast(in), static_cast(out), s_input, - s_output, dim, total); + generic_sum_dim_kernel<<>>(static_cast(in), + out, s_input, s_output, dim, + total); CHECK_CUDA(cudaGetLastError()); } -void launch_sum_dim(void *out, void *in, const StrideAndSize &s_input, - const StrideAndSize &s_output, int64_t dim) { +template +void launch_sum_dim_accum(float *out, const void *in, + const StrideAndSize &s_input, + const StrideAndSize &s_output, int64_t dim) { if (s_input.rank <= 3) { const StrideAndSize opt_input = padded_rank3(s_input); const StrideAndSize opt_output = padded_rank3(s_output); if (dim == 0) { - launch_sum_dim0(out, in, opt_input, opt_output); + launch_sum_dim0(out, in, opt_input, opt_output); return; } if (dim == 1) { - launch_sum_dim1(out, in, opt_input, opt_output); + launch_sum_dim1(out, in, opt_input, opt_output); return; } - launch_sum_dim2(out, in, opt_input, opt_output); + launch_sum_dim2(out, in, opt_input, opt_output); + return; + } + + launch_generic_sum_dim(out, in, s_input, s_output, dim); +} + +} // namespace + +void launch_sum_dim(void *out, const void *in, DataType input_dtype, + const StrideAndSize &s_input, + const StrideAndSize &s_output, int64_t dim) { + const size_t output_total = shape_numel(s_output.size, s_output.rank); + if (output_total == 0) { return; } - launch_generic_sum_dim(out, in, s_input, s_output, dim); + if (input_dtype == DataType::f32) { + launch_sum_dim_accum(static_cast(out), in, s_input, + s_output, dim); + return; + } + + ASSERT(input_dtype == DataType::f16, + fmt::format("Unsupported dtype {}", get_name(input_dtype))); + + launch_sum_dim_accum<__half>(static_cast(out), in, s_input, + s_output, dim); } } // namespace smollnet From 46ecbec3da6faa053007e8bf057b69bf8fb06d64 Mon Sep 17 00:00:00 2001 From: Jacob Domagala Date: Thu, 13 Nov 2025 12:05:00 +0100 Subject: [PATCH 06/13] [#81]: Dispatch CUDA math kernels by dtype --- src/autograd.cpp | 30 ++- src/kernels.cu | 678 ++++++++++++++++++++++++++++------------------- src/sgd.cpp | 3 +- 3 files changed, 425 insertions(+), 286 deletions(-) diff --git a/src/autograd.cpp b/src/autograd.cpp index ff547e9..9f23906 100644 --- a/src/autograd.cpp +++ b/src/autograd.cpp @@ -1,4 +1,5 @@ #include "autograd.hpp" +#include "dtype_utils.hpp" #include "helpers.hpp" #include "kernels.cuh" @@ -116,7 +117,7 @@ SubFunction::backward(const std::vector &grad_outputs) { } } - launch_negative(grad.data(), grad.numel()); + launch_negative(grad.data(), grad.dtype(), grad.numel()); grad_inputs[1] = grad; } @@ -187,7 +188,7 @@ DivFunction::backward(const std::vector &grad_outputs) { if (needs_input_grad[1]) { Tensor grad = (grad_outputs[0] * inputs[0]) / (inputs[1] * inputs[1]); grad = reduce_broadcast_gradient(grad, inputs[1]); - launch_negative(grad.data(), grad.numel()); + launch_negative(grad.data(), grad.dtype(), grad.numel()); grad_inputs[1] = grad; } @@ -239,7 +240,7 @@ ReLUFunction::backward(const std::vector &grad_outputs) { if (needs_input_grad[0]) { gi[0] = create_grad_tensor(inputs[0]); launch_relu_grad(gi[0].data(), grad_outputs[0].data(), inputs[0].data(), - gi[0].numel()); + gi[0].dtype(), gi[0].numel()); } return gi; } @@ -257,7 +258,7 @@ GeLUFunction::backward(const std::vector &grad_outputs) { if (needs_input_grad[0]) { gi[0] = create_grad_tensor(inputs[0]); launch_gelu_grad(gi[0].data(), grad_outputs[0].data(), inputs[0].data(), - gi[0].numel()); + gi[0].dtype(), gi[0].numel()); } return gi; } @@ -278,7 +279,8 @@ TanhFunction::backward(const std::vector &grad_outputs) { if (needs_input_grad[0]) { auto grad_input = create_grad_tensor(inputs[0]); launch_tanh_grad(grad_input.data(), grad_outputs.front().data(), - inputs[0].data(), grad_input.numel()); + inputs[0].data(), grad_input.dtype(), + grad_input.numel()); grad_inputs[0] = grad_input; } @@ -302,7 +304,8 @@ SigmoidFunction::backward(const std::vector &grad_outputs) { auto grad_input = create_grad_tensor(inputs[0]); launch_sigmoid_grad(grad_input.data(), grad_outputs.front().data(), - inputs[0].data(), grad_input.numel()); + inputs[0].data(), grad_input.dtype(), + grad_input.numel()); grad_inputs[0] = grad_input; } @@ -341,17 +344,20 @@ MseFunction::MseFunction(const Tensor &pred, const Tensor &tgt) : N(pred.numel() std::vector MseFunction::backward(const std::vector &grad_outputs) { ASSERT(grad_outputs.size() == 1, "MSE backward expects 1 grad_output (scalar)"); - float c = - *static_cast(grad_outputs[0].cpu().data()) * (2.f / static_cast(N)); + Tensor grad_output_cpu = grad_outputs[0].cpu(); + float c = load_scalar(grad_output_cpu.data(), grad_output_cpu.dtype(), 0) * + (2.f / static_cast(N)); std::vector gi(2); if (needs_input_grad[0]) { gi[0] = create_grad_tensor(inputs[0]); - launch_mse_grad(gi[0].data(), inputs[0].data(), inputs[1].data(), c, N); + launch_mse_grad(gi[0].data(), inputs[0].data(), inputs[1].data(), + gi[0].dtype(), c, N); } if (needs_input_grad[1]) { gi[1] = create_grad_tensor(inputs[1]); - launch_mse_grad(gi[1].data(), inputs[0].data(), inputs[1].data(), -c, N); + launch_mse_grad(gi[1].data(), inputs[0].data(), inputs[1].data(), + gi[1].dtype(), -c, N); } return gi; } @@ -392,8 +398,8 @@ LayerNormFunction::backward(const std::vector &grad_outputs) { hat_x.device(), true); launch_layer_norm_grad(dx.data(), hat_x.data(), delta.data(), variance.data(), - sum_delta.data(), sum_dh.data(), hat_x.size(0), - hat_x.size(1)); + sum_delta.data(), sum_dh.data(), dx.dtype(), + variance.dtype(), hat_x.size(0), hat_x.size(1)); /* ---------- gradient w.r.t. γ (scale) ---------- */ // sum over batch → shape [F,1] diff --git a/src/kernels.cu b/src/kernels.cu index 0ac4968..9450085 100644 --- a/src/kernels.cu +++ b/src/kernels.cu @@ -1,252 +1,168 @@ +#include "dtype_utils.hpp" #include "helpers.hpp" #include "kernels.cuh" -#include -#include - #include +#include namespace smollnet { -__global__ void matmul_kernel(float *__restrict__ C, - const float *__restrict__ A, - const float *__restrict__ B, + +namespace { + +template +__global__ void matmul_kernel(T *__restrict__ C, const T *__restrict__ A, + const T *__restrict__ B, const StrideInfo strides, const SizeInfo sizes, const int tile_width) { - const int col = blockIdx.x * blockDim.x + threadIdx.x; // N‑index - const int row = blockIdx.y * blockDim.y + threadIdx.y; // M‑index + const int col = blockIdx.x * blockDim.x + threadIdx.x; + const int row = blockIdx.y * blockDim.y + threadIdx.y; const int M = strides.output_size[0]; const int N = strides.output_size[1]; - const int K = sizes.a_size[1]; // = sizes.b_size[0] + const int K = sizes.a_size[1]; const bool in_bounds = (row < M) && (col < N); extern __shared__ float s_mem[]; - float *As = s_mem; // tile from A (M×K) - float *Bs = s_mem + tile_width * tile_width; // tile from B (K×N) + float *As = s_mem; + float *Bs = s_mem + tile_width * tile_width; float acc = 0.0f; const int num_tiles = (K + tile_width - 1) / tile_width; for (int t = 0; t < num_tiles; ++t) { - const int a_col = t * tile_width + threadIdx.x; // K‑index into A - const int b_row = t * tile_width + threadIdx.y; // K‑index into B + const int a_col = t * tile_width + threadIdx.x; + const int b_row = t * tile_width + threadIdx.y; - // Load current tiles into shared memory, zero‑padding out‑of‑range - // elements. + const int64_t a_offset = + row * strides.a_stride[0] + a_col * strides.a_stride[1]; As[threadIdx.y * tile_width + threadIdx.x] = - (row < M && a_col < K) ? A[row * K + a_col] : 0.0f; + (row < M && a_col < K) ? scalar_to_float(A[a_offset]) : 0.0f; + const int64_t b_offset = + b_row * strides.b_stride[0] + col * strides.b_stride[1]; Bs[threadIdx.y * tile_width + threadIdx.x] = - (b_row < K && col < N) ? B[b_row * N + col] : 0.0f; + (b_row < K && col < N) ? scalar_to_float(B[b_offset]) : 0.0f; __syncthreads(); - // Multiply–accumulate over the valid fragment length. const int elems = min(tile_width, K - t * tile_width); #pragma unroll - for (int e = 0; e < elems; ++e) + for (int e = 0; e < elems; ++e) { acc += As[threadIdx.y * tile_width + e] * Bs[e * tile_width + threadIdx.x]; + } __syncthreads(); } - if (in_bounds) - C[row * N + col] = acc; + if (in_bounds) { + C[row * N + col] = scalar_from_float(acc); + } } -void launch_matmul(void *out, void *left, void *right, - const StrideInfo &strides, const SizeInfo &sizes, - size_t total) { - constexpr int TILE = 16; - dim3 block(TILE, TILE); +template +__global__ void unary_kernel(T *out, const T *in, size_t total, Op op) { + const auto idx = threadIdx.x + blockDim.x * blockIdx.x; - const int M = strides.output_size[0]; // rows of C - const int N = strides.output_size[1]; // cols of C - - dim3 grid((N + TILE - 1) / TILE, // x‑dim ← N - (M + TILE - 1) / TILE); // y‑dim ← M - - size_t smem_bytes = 2 * TILE * TILE * sizeof(float); - - matmul_kernel<<>>( - static_cast(out), static_cast(left), - static_cast(right), strides, sizes, TILE); - - CHECK_CUDA(cudaGetLastError()); + if (idx < total) { + out[idx] = scalar_from_float(op(scalar_to_float(in[idx]))); + } } -__global__ void relu_kernel(float *out, float *in, size_t total) { - auto idx = threadIdx.x + blockDim.x * blockIdx.x; - - if (idx < total) - out[idx] = fmaxf(in[idx], 0.0f); -} +struct ReluOp { + __device__ float operator()(float x) const { return fmaxf(x, 0.0f); } +}; -void launch_relu(void *out, void *in, size_t total) { +struct GeluOp { + __device__ float operator()(float x) const { + constexpr float sqrt_2_over_pi = 0.7978845608f; + return 0.5f * x * + (1.0f + tanhf(sqrt_2_over_pi * (x + 0.044715f * x * x * x))); + } +}; - int block = 256; - int grid = (total + block - 1) / block; +struct TanhOp { + __device__ float operator()(float x) const { return tanhf(x); } +}; - relu_kernel<<>>(static_cast(out), - static_cast(in), total); - CHECK_CUDA(cudaGetLastError()); -} +struct SigmoidOp { + __device__ float operator()(float x) const { + return 1.0f / (1.0f + expf(-x)); + } +}; -__global__ void relu_grad_kernel(float *out, float *grad_out, float *in, +template +__global__ void relu_grad_kernel(T *out, const T *grad_out, const T *in, size_t total) { - auto idx = threadIdx.x + blockDim.x * blockIdx.x; - - if (idx < total) - out[idx] = (in[idx] > 0.0f) ? grad_out[idx] : 0.0f; -} - -void launch_relu_grad(void *out, void *grad_out, void *in, size_t total) { - int block = 256; - int grid = (total + block - 1) / block; - - relu_grad_kernel<<>>(static_cast(out), - static_cast(grad_out), - static_cast(in), total); - CHECK_CUDA(cudaGetLastError()); -} - -__global__ void gelu_kernel(float *out, float *in, size_t total) { - auto idx = threadIdx.x + blockDim.x * blockIdx.x; + const auto idx = threadIdx.x + blockDim.x * blockIdx.x; if (idx < total) { - constexpr float sqrt_2_over_pi = 0.7978845608f; - out[idx] = - 0.5f * in[idx] * - (1.0f + tanhf(sqrt_2_over_pi * - (in[idx] + 0.044715f * in[idx] * in[idx] * in[idx]))); + const float input = scalar_to_float(in[idx]); + const float grad = input > 0.0f ? scalar_to_float(grad_out[idx]) : 0.0f; + out[idx] = scalar_from_float(grad); } } -void launch_gelu(void *out, void *in, size_t total) { - - int block = 256; - int grid = (total + block - 1) / block; - - gelu_kernel<<>>(static_cast(out), - static_cast(in), total); - CHECK_CUDA(cudaGetLastError()); -} - -__global__ void gelu_grad_kernel(float *out, float *grad_out, float *in, +template +__global__ void gelu_grad_kernel(T *out, const T *grad_out, const T *in, size_t total) { - auto idx = threadIdx.x + blockDim.x * blockIdx.x; + const auto idx = threadIdx.x + blockDim.x * blockIdx.x; if (idx < total) { - // const float a = std::sqrt(2.0f / M_PI); constexpr float a = 0.7978845608f; constexpr float b = 0.044715f; - float x3 = in[idx] * in[idx] * in[idx]; - float h = in[idx] + b * x3; - float tanh_ax = tanhf(a * h); - float sech2 = 1.0f - tanh_ax * tanh_ax; - float h_prime = 1.0f + 3.0f * b * in[idx] * in[idx]; - - float g = 0.5f * (1.0f + tanh_ax) + 0.5f * in[idx] * sech2 * a * h_prime; - out[idx] = grad_out[idx] * g; + const float x = scalar_to_float(in[idx]); + const float x3 = x * x * x; + const float h = x + b * x3; + const float tanh_ax = tanhf(a * h); + const float sech2 = 1.0f - tanh_ax * tanh_ax; + const float h_prime = 1.0f + 3.0f * b * x * x; + + const float g = + 0.5f * (1.0f + tanh_ax) + 0.5f * x * sech2 * a * h_prime; + out[idx] = scalar_from_float(scalar_to_float(grad_out[idx]) * g); } } -void launch_gelu_grad(void *out, void *grad_out, void *in, size_t total) { - int block = 256; - int grid = (total + block - 1) / block; - - gelu_grad_kernel<<>>(static_cast(out), - static_cast(grad_out), - static_cast(in), total); - CHECK_CUDA(cudaGetLastError()); -} - -__global__ void tanh_kernel(float *out, float *in, size_t total) { - auto idx = threadIdx.x + blockDim.x * blockIdx.x; - - if (idx < total) - out[idx] = tanhf(in[idx]); -} - -void launch_tanh(void *out, void *in, size_t total) { - - int block = 256; - int grid = (total + block - 1) / block; - - tanh_kernel<<>>(static_cast(out), - static_cast(in), total); - CHECK_CUDA(cudaGetLastError()); -} - -__global__ void tanh_grad_kernel(float *out, float *grad_out, float *in, +template +__global__ void tanh_grad_kernel(T *out, const T *grad_out, const T *in, size_t total) { - auto idx = threadIdx.x + blockDim.x * blockIdx.x; - - if (idx < total) - out[idx] = grad_out[idx] * (1.0f - in[idx] * in[idx]); -} - -void launch_tanh_grad(void *out, void *grad_out, void *in, size_t total) { + const auto idx = threadIdx.x + blockDim.x * blockIdx.x; - int block = 256; - int grid = (total + block - 1) / block; - - tanh_grad_kernel<<>>(static_cast(out), - static_cast(grad_out), - static_cast(in), total); - CHECK_CUDA(cudaGetLastError()); -} - -__global__ void sigmoid_kernel(float *output, float *input, int size) { - int idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < size) { - output[idx] = - 1.0f / (1.0f + expf(-input[idx])); // Apply sigmoid to each element + if (idx < total) { + const float x = scalar_to_float(in[idx]); + const float grad = scalar_to_float(grad_out[idx]) * (1.0f - x * x); + out[idx] = scalar_from_float(grad); } } -void launch_sigmoid(void *out, void *in, size_t total) { - - int block = 256; - int grid = (total + block - 1) / block; - - sigmoid_kernel<<>>(static_cast(out), - static_cast(in), total); - CHECK_CUDA(cudaGetLastError()); -} - -__global__ void sigmoid_grad_kernel(float *output, float *grad_out, - float *input, int size) { - int idx = blockIdx.x * blockDim.x + threadIdx.x; +template +__global__ void sigmoid_grad_kernel(T *output, const T *grad_out, + const T *input, int size) { + const int idx = blockIdx.x * blockDim.x + threadIdx.x; if (idx < size) { - output[idx] = grad_out[idx] * input[idx] * (1.0f - input[idx]); + const float x = scalar_to_float(input[idx]); + const float grad = scalar_to_float(grad_out[idx]) * x * (1.0f - x); + output[idx] = scalar_from_float(grad); } } -void launch_sigmoid_grad(void *out, void *grad_out, void *in, size_t total) { - - int block = 256; - int grid = (total + block - 1) / block; - - sigmoid_grad_kernel<<>>(static_cast(out), - static_cast(grad_out), - static_cast(in), total); - CHECK_CUDA(cudaGetLastError()); -} - -__global__ void mse_kernel(float *out, const float *__restrict__ pred, - const float *__restrict__ target, std::size_t n) { +template +__global__ void mse_partial_kernel(float *partial, + const InT *__restrict__ pred, + const InT *__restrict__ target, + std::size_t n) { __shared__ float sMem[32]; std::size_t idx = blockIdx.x * blockDim.x + threadIdx.x; - std::size_t stride = blockDim.x * gridDim.x; + const std::size_t stride = blockDim.x * gridDim.x; float local_sum = 0.0f; for (; idx < n; idx += stride) { - float diff = pred[idx] - target[idx]; + const float diff = + scalar_to_float(pred[idx]) - scalar_to_float(target[idx]); local_sum += diff * diff; } @@ -255,9 +171,9 @@ __global__ void mse_kernel(float *out, const float *__restrict__ pred, local_sum += __shfl_down_sync(0xffffffff, local_sum, off); } - int lane = threadIdx.x & 31; - int warp = threadIdx.x >> 5; - int num_warps = (blockDim.x + 31) >> 5; + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + const int num_warps = (blockDim.x + 31) >> 5; if (lane == 0) { sMem[warp] = local_sum; @@ -277,148 +193,364 @@ __global__ void mse_kernel(float *out, const float *__restrict__ pred, } if (threadIdx.x == 0) { - atomicAdd(out, block_sum / static_cast(n)); + partial[blockIdx.x] = block_sum; } } -void launch_mse(void *out, void *pred, void *target, size_t total) { - constexpr int BLOCK_SIZE = 256; - int grid = (total + BLOCK_SIZE - 1) / BLOCK_SIZE; +template +__global__ void mse_finalize_kernel(OutT *out, const float *partial, + std::size_t num_partials, + std::size_t n) { + float acc = 0.0f; + for (std::size_t idx = threadIdx.x; idx < num_partials; idx += blockDim.x) { + acc += partial[idx]; + } - mse_kernel<<>>(static_cast(out), - static_cast(pred), - static_cast(target), total); +#pragma unroll + for (int off = 16; off > 0; off >>= 1) { + acc += __shfl_down_sync(0xffffffff, acc, off); + } - CHECK_CUDA(cudaGetLastError()); + __shared__ float sMem[32]; + const int lane = threadIdx.x & 31; + const int warp = threadIdx.x >> 5; + const int num_warps = (blockDim.x + 31) >> 5; + + if (lane == 0) { + sMem[warp] = acc; + } + + __syncthreads(); + + if (warp == 0) { + acc = (lane < num_warps) ? sMem[lane] : 0.0f; +#pragma unroll + for (int off = 16; off > 0; off >>= 1) { + acc += __shfl_down_sync(0xffffffff, acc, off); + } + } + + if (threadIdx.x == 0) { + out[0] = scalar_from_float(acc / static_cast(n)); + } } -__global__ void sgd_kernel(float *w, const float *grad, float lr, + +template +__global__ void sgd_kernel(ParamT *w, const GradT *grad, float lr, size_t total) { - size_t idx = blockIdx.x * blockDim.x + threadIdx.x; + const size_t idx = blockIdx.x * blockDim.x + threadIdx.x; if (idx < total) { - w[idx] -= lr * grad[idx]; + const float updated = + scalar_to_float(w[idx]) - lr * scalar_to_float(grad[idx]); + w[idx] = scalar_from_float(updated); } } -void launch_sgd_update(void *p, void *g, float lr, size_t total) { - dim3 block = 256; - dim3 grid = (total + block.x - 1) / block.x; - sgd_kernel<<>>(static_cast(p), - static_cast(g), lr, total); +template +__global__ void mse_grad_kernel(T *g, const T *p, const T *t, float coeff, + size_t n) { + const size_t idx = blockIdx.x * blockDim.x + threadIdx.x; + if (idx < n) { + const float grad = + coeff * (scalar_to_float(p[idx]) - scalar_to_float(t[idx])); + g[idx] = scalar_from_float(grad); + } +} + +template +__global__ void mean_2d_kernel(OutT *out, const InT *in, size_t d0, + size_t d1) { + const auto idx = threadIdx.x + blockDim.x * blockIdx.x; + + if (idx >= d0) { + return; + } + + float acc = 0.0f; + for (size_t i = 0; i < d1; ++i) { + acc += scalar_to_float(in[idx * d1 + i]); + } + + out[idx] = scalar_from_float(acc / static_cast(d1)); +} + +template +__global__ void layer_norm_kernel(DataT *out, const DataT *features, + const StatsT *mean, const StatsT *variance, + const DataT *gamma, const DataT *beta, + size_t batch_size, size_t num_features) { + const auto idx = threadIdx.x + blockDim.x * blockIdx.x; + const auto total = batch_size * num_features; + + if (idx >= total) { + return; + } + + const size_t batch_num = idx / num_features; + + constexpr float epsilon = 1e-5f; + const float normalized = + (scalar_to_float(features[idx]) - scalar_to_float(mean[batch_num])) / + sqrtf(scalar_to_float(variance[batch_num]) + epsilon); + + const float value = scalar_to_float(gamma[batch_num]) * normalized + + scalar_to_float(beta[batch_num]); + out[idx] = scalar_from_float(value); +} + +template +__global__ void layer_norm_grad_kernel( + DataT *out_grad, const DataT *normalized_input, + const DataT *scaled_gradient, const StatsT *variance, + const StatsT *summed_scale, const StatsT *summed_scaled_input, + size_t batch_size, size_t num_features) { + const size_t idx = threadIdx.x + blockDim.x * blockIdx.x; + const size_t total = batch_size * num_features; + if (idx >= total) { + return; + } + + const size_t row = idx / num_features; + + constexpr float eps = 1e-5f; + const float inv_std = rsqrtf(scalar_to_float(variance[row]) + eps); + const float m1 = scalar_to_float(summed_scale[row]) / num_features; + const float m2 = scalar_to_float(summed_scaled_input[row]) / num_features; + + const float hat_x = scalar_to_float(normalized_input[idx]); + const float delta = scalar_to_float(scaled_gradient[idx]); + + const float res = inv_std * (delta - m1 - hat_x * m2); + out_grad[idx] = scalar_from_float(res); +} + +template +void launch_unary(void *out, const void *in, DataType dtype, size_t total, + Op op) { + const int block = 256; + const int grid = (total + block - 1) / block; + + dispatch_float_dtype(dtype, [&]() { + unary_kernel<<>>(static_cast(out), + static_cast(in), total, + op); + }); + + CHECK_CUDA(cudaGetLastError()); +} + +} // namespace + +void launch_matmul(void *out, const void *left, const void *right, + DataType dtype, const StrideInfo &strides, + const SizeInfo &sizes, size_t /*total*/) { + constexpr int TILE = 16; + dim3 block(TILE, TILE); + + const int M = strides.output_size[0]; + const int N = strides.output_size[1]; + + dim3 grid((N + TILE - 1) / TILE, (M + TILE - 1) / TILE); + + const size_t smem_bytes = 2 * TILE * TILE * sizeof(float); + + dispatch_float_dtype(dtype, [&]() { + matmul_kernel<<>>( + static_cast(out), static_cast(left), + static_cast(right), strides, sizes, TILE); + }); + CHECK_CUDA(cudaGetLastError()); } -__global__ void mse_grad_kernel(float *g, const float *p, const float *t, - float coeff, size_t n) { - size_t idx = blockIdx.x * blockDim.x + threadIdx.x; - if (idx < n) - g[idx] = coeff * (p[idx] - t[idx]); +void launch_relu(void *out, const void *in, DataType dtype, size_t total) { + launch_unary(out, in, dtype, total, ReluOp{}); } -void launch_mse_grad(void *grad, void *pred, void *target, float coeff, - size_t total) { - int block = 256; - int grid = (total + block - 1) / block; - mse_grad_kernel<<>>(static_cast(grad), - static_cast(pred), - static_cast(target), coeff, total); +void launch_relu_grad(void *out, const void *grad_out, const void *in, + DataType dtype, size_t total) { + const int block = 256; + const int grid = (total + block - 1) / block; + + dispatch_float_dtype(dtype, [&]() { + relu_grad_kernel<<>>( + static_cast(out), static_cast(grad_out), + static_cast(in), total); + }); + CHECK_CUDA(cudaGetLastError()); } -__global__ void mean_2d_kernel(float *out, float *in, size_t d0, size_t d1) { - auto idx = threadIdx.x + blockDim.x * blockIdx.x; +void launch_gelu(void *out, const void *in, DataType dtype, size_t total) { + launch_unary(out, in, dtype, total, GeluOp{}); +} - if (idx >= d0) - return; +void launch_gelu_grad(void *out, const void *grad_out, const void *in, + DataType dtype, size_t total) { + const int block = 256; + const int grid = (total + block - 1) / block; - float acc = 0.0f; - for (int i = 0; i < d1; ++i) { - acc += in[idx * d1 + i]; - } + dispatch_float_dtype(dtype, [&]() { + gelu_grad_kernel<<>>( + static_cast(out), static_cast(grad_out), + static_cast(in), total); + }); - acc /= d1; + CHECK_CUDA(cudaGetLastError()); +} - out[idx] = acc; +void launch_tanh(void *out, const void *in, DataType dtype, size_t total) { + launch_unary(out, in, dtype, total, TanhOp{}); } -void launch_mean_2d(void *out, void *in, size_t d0, size_t d1) { - dim3 block = 256; - dim3 grid = (block.x + d0 * d1 - 1) / block.x; +void launch_tanh_grad(void *out, const void *grad_out, const void *in, + DataType dtype, size_t total) { + const int block = 256; + const int grid = (total + block - 1) / block; - mean_2d_kernel<<>>(static_cast(out), - static_cast(in), d0, d1); + dispatch_float_dtype(dtype, [&]() { + tanh_grad_kernel<<>>( + static_cast(out), static_cast(grad_out), + static_cast(in), total); + }); + + CHECK_CUDA(cudaGetLastError()); } -__global__ void layer_norm_kernel(float *out, float *features, float *mean, - float *variance, float *gamma, float *beta, - size_t batch_size, size_t num_features) { - auto idx = threadIdx.x + blockDim.x * blockIdx.x; - const auto total = batch_size * num_features; +void launch_sigmoid(void *out, const void *in, DataType dtype, size_t total) { + launch_unary(out, in, dtype, total, SigmoidOp{}); +} - if (idx >= total) - return; +void launch_sigmoid_grad(void *out, const void *grad_out, const void *in, + DataType dtype, size_t total) { + const int block = 256; + const int grid = (total + block - 1) / block; - int batch_num = idx / num_features; + dispatch_float_dtype(dtype, [&]() { + sigmoid_grad_kernel<<>>( + static_cast(out), static_cast(grad_out), + static_cast(in), static_cast(total)); + }); - constexpr float epsilon = 1e-5f; - float normalized = - (features[idx] - mean[batch_num]) / sqrtf(variance[batch_num] + epsilon); + CHECK_CUDA(cudaGetLastError()); +} - out[idx] = gamma[batch_num] * normalized + beta[batch_num]; +void launch_mse(void *out, DataType out_dtype, const void *pred, + const void *target, DataType input_dtype, size_t total) { + constexpr int BLOCK_SIZE = 256; + const int grid = static_cast((total + BLOCK_SIZE - 1) / BLOCK_SIZE); + + float *partials = nullptr; + CHECK_CUDA(cudaMalloc(&partials, grid * sizeof(float))); + + dispatch_float_dtype(input_dtype, [&]() { + mse_partial_kernel<<>>( + partials, static_cast(pred), + static_cast(target), total); + }); + CHECK_CUDA(cudaGetLastError()); + + dispatch_float_dtype(out_dtype, [&]() { + mse_finalize_kernel<<<1, BLOCK_SIZE>>>( + static_cast(out), partials, static_cast(grid), + total); + }); + CHECK_CUDA(cudaGetLastError()); + + CHECK_CUDA(cudaFree(partials)); } -void launch_layer_norm(void *out, void *features, void *mean, void *variance, - void *gamma, void *beta, size_t batch_size, - size_t num_features) { +void launch_sgd_update(void *p, const void *g, DataType param_dtype, + DataType grad_dtype, float lr, size_t total) { dim3 block = 256; - size_t total = batch_size * num_features; - dim3 grid = (block.x + total - 1) / block.x; + dim3 grid = (total + block.x - 1) / block.x; + + dispatch_float_dtype(param_dtype, [&]() { + dispatch_float_dtype(grad_dtype, [&]() { + sgd_kernel<<>>( + static_cast(p), static_cast(g), lr, total); + }); + }); + + CHECK_CUDA(cudaGetLastError()); +} + +void launch_mse_grad(void *grad, const void *pred, const void *target, + DataType dtype, float coeff, size_t total) { + const int block = 256; + const int grid = (total + block - 1) / block; - layer_norm_kernel<<>>( - static_cast(out), static_cast(features), - static_cast(mean), static_cast(variance), - static_cast(gamma), static_cast(beta), batch_size, - num_features); + dispatch_float_dtype(dtype, [&]() { + mse_grad_kernel<<>>( + static_cast(grad), static_cast(pred), + static_cast(target), coeff, total); + }); + + CHECK_CUDA(cudaGetLastError()); } -__global__ void layer_norm_grad_kernel(float *out_grad, - const float *normalized_input, - const float *scaled_gradient, - const float *variance, - const float *summed_scale, - const float *summed_scaled_input, - size_t batch_size, size_t num_features) { - const size_t idx = threadIdx.x + blockDim.x * blockIdx.x; - const size_t total = batch_size * num_features; - if (idx >= total) - return; +void launch_mean_2d(void *out, DataType out_dtype, const void *in, + DataType in_dtype, size_t d0, size_t d1) { + dim3 block = 256; + dim3 grid = (block.x + d0 - 1) / block.x; - const size_t row = idx / num_features; + dispatch_float_dtype(out_dtype, [&]() { + dispatch_float_dtype(in_dtype, [&]() { + mean_2d_kernel<<>>( + static_cast(out), static_cast(in), d0, d1); + }); + }); - constexpr float eps = 1e-5f; - const float inv_std = rsqrtf(variance[row] + eps); // per-sample variance - const float m1 = summed_scale[row] / num_features; // Σδ / D - const float m2 = summed_scaled_input[row] / num_features; // Σδ·ẋ / D + CHECK_CUDA(cudaGetLastError()); +} + +void launch_layer_norm(void *out, const void *features, const void *mean, + const void *variance, const void *gamma, + const void *beta, DataType data_dtype, + DataType stats_dtype, size_t batch_size, + size_t num_features) { + dim3 block = 256; + const size_t total = batch_size * num_features; + dim3 grid = (block.x + total - 1) / block.x; - const float hat_x = normalized_input[idx]; - const float delta = scaled_gradient[idx]; // δ = dy * γ + dispatch_float_dtype(data_dtype, [&]() { + dispatch_float_dtype(stats_dtype, [&]() { + layer_norm_kernel<<>>( + static_cast(out), static_cast(features), + static_cast(mean), + static_cast(variance), + static_cast(gamma), static_cast(beta), + batch_size, num_features); + }); + }); - const float res = inv_std * (delta - m1 - hat_x * m2); // ∂L/∂x - out_grad[idx] = res; + CHECK_CUDA(cudaGetLastError()); } -void launch_layer_norm_grad(void *out, void *normalized_input, - void *scaled_gradient, void *variance, - void *summed_scale, void *summed_scaled_input, +void launch_layer_norm_grad(void *out, const void *normalized_input, + const void *scaled_gradient, const void *variance, + const void *summed_scale, + const void *summed_scaled_input, + DataType data_dtype, DataType stats_dtype, size_t batch_size, size_t num_features) { - dim3 block = 256; - size_t total = batch_size * num_features; + const size_t total = batch_size * num_features; dim3 grid = (block.x + total - 1) / block.x; - layer_norm_grad_kernel<<>>( - static_cast(out), static_cast(normalized_input), - static_cast(scaled_gradient), static_cast(variance), - static_cast(summed_scale), - static_cast(summed_scaled_input), batch_size, num_features); + + dispatch_float_dtype(data_dtype, [&]() { + dispatch_float_dtype(stats_dtype, [&]() { + layer_norm_grad_kernel<<>>( + static_cast(out), + static_cast(normalized_input), + static_cast(scaled_gradient), + static_cast(variance), + static_cast(summed_scale), + static_cast(summed_scaled_input), batch_size, + num_features); + }); + }); + + CHECK_CUDA(cudaGetLastError()); } } // namespace smollnet diff --git a/src/sgd.cpp b/src/sgd.cpp index 7c99b8a..8da5a09 100644 --- a/src/sgd.cpp +++ b/src/sgd.cpp @@ -17,7 +17,8 @@ void SGD::step() const { grad.size(dim))); } - launch_sgd_update(p.data(), grad.data(), lr_, p.numel()); + launch_sgd_update(p.data(), grad.data(), p.dtype(), grad.dtype(), lr_, + p.numel()); } } From 61d2928e906f75373855a71b72bee7f89407b338 Mon Sep 17 00:00:00 2001 From: Jacob Domagala Date: Thu, 13 Nov 2025 12:06:00 +0100 Subject: [PATCH 07/13] [#81]: Accumulate layer norm stats in fp32 --- src/layer_norm.cpp | 28 ++++++++----- src/welford.cu | 102 ++++++++++++++++++++++++++------------------- 2 files changed, 77 insertions(+), 53 deletions(-) diff --git a/src/layer_norm.cpp b/src/layer_norm.cpp index b878b2f..63c4fbc 100644 --- a/src/layer_norm.cpp +++ b/src/layer_norm.cpp @@ -1,6 +1,7 @@ #include "layer_norm.hpp" -#include "kernels.cuh" #include "autograd.hpp" +#include "dtype_utils.hpp" +#include "kernels.cuh" #include @@ -17,21 +18,28 @@ Tensor LayerNorm::compute(const Tensor &t) { bias = zeros({t.size(1), 1}, t.dtype(), t.device(), true); } - auto mean = zeros({t.size(0), 1}, t.dtype(), t.device()); - launch_mean_2d(mean.data(), t.data(), t.size(0), t.size(1)); + const DataType stats_dtype = accumulation_dtype(t.dtype()); + + auto mean = zeros({t.size(0), 1}, stats_dtype, t.device()); + launch_mean_2d(mean.data(), mean.dtype(), t.data(), t.dtype(), t.size(0), + t.size(1)); - auto variance = zeros({t.size(0), 1}, t.dtype(), t.device()); - launch_welford(t.data(), variance.data(), t.size(1), t.size(0), 0, WelfordType::PopulationVariance); + auto variance = zeros({t.size(0), 1}, stats_dtype, t.device()); + launch_welford(t.data(), t.dtype(), variance.data(), variance.dtype(), + t.size(1), t.size(0), 0, WelfordType::PopulationVariance); - auto normalized = zeros(t.dims().data(), t.ndims(), t.dtype(), t.device(), t.requires_grad()); + auto normalized = zeros(t.dims().data(), t.ndims(), t.dtype(), t.device(), + t.requires_grad()); launch_layer_norm(normalized.data(), t.data(), mean.data(), variance.data(), - weights.data(), bias.data(), t.size(0), t.size(1)); + weights.data(), bias.data(), normalized.dtype(), + mean.dtype(), t.size(0), t.size(1)); - if(normalized.requires_grad()) { - auto* meta = normalized.autograd(); + if (normalized.requires_grad()) { + auto *meta = normalized.autograd(); meta->is_leaf = false; - meta->grad_fn = std::make_shared(mean, variance, normalized, t, weights, bias); + meta->grad_fn = std::make_shared( + mean, variance, normalized, t, weights, bias); } return normalized; diff --git a/src/welford.cu b/src/welford.cu index 12ce94a..4b2125d 100644 --- a/src/welford.cu +++ b/src/welford.cu @@ -1,6 +1,7 @@ #include #include +#include "dtype_utils.hpp" #include "kernels.cuh" #include "welford_internal.inl" namespace smollnet { @@ -40,8 +41,8 @@ __device__ __forceinline__ WelfordData merge(const WelfordData &a, return out; } -template -__global__ void welford_row_first_pass(const float *__restrict__ in, +template +__global__ void welford_row_first_pass(const InT *__restrict__ in, WelfordData *__restrict__ out, const size_t num_features, const size_t num_rows) { @@ -53,7 +54,7 @@ __global__ void welford_row_first_pass(const float *__restrict__ in, uint32_t base_idx = feature + batch * num_features * CHUNK_SIZE; - WelfordData localData; + WelfordData localData{}; __shared__ WelfordData sMem[BLOCK_DIM / 32]; @@ -67,7 +68,7 @@ __global__ void welford_row_first_pass(const float *__restrict__ in, if (row >= num_rows) return; - float v = is_valid ? in[idx] : 0.0f; + float v = is_valid ? scalar_to_float(in[idx]) : 0.0f; update(localData, v, is_valid); #pragma unroll @@ -102,10 +103,10 @@ __global__ void welford_row_first_pass(const float *__restrict__ in, } } -template +template __global__ void welford_row_second_pass(const WelfordData *__restrict__ in, - float *__restrict__ out, const int32_t num_elems, + OutT *__restrict__ out, const int32_t num_elems, const int32_t num_iter, WelfordType type) { const auto col = threadIdx.x; const auto row = blockIdx.x; @@ -154,22 +155,24 @@ welford_row_second_pass(const WelfordData *__restrict__ in, if (threadIdx.x == 0) { switch (type) { case WelfordType::Mean: { - out[blockIdx.x] = final_result.mean; + out[blockIdx.x] = scalar_from_float(final_result.mean); } break; case WelfordType::PopulationVariance: { - out[blockIdx.x] = final_result.M2 / final_result.count; + out[blockIdx.x] = + scalar_from_float(final_result.M2 / final_result.count); } break; case WelfordType::SampleVariance: { - out[blockIdx.x] = final_result.M2 / (final_result.count - 1); + out[blockIdx.x] = + scalar_from_float(final_result.M2 / (final_result.count - 1)); } break; } } } -template -__global__ void welford_column_first_pass(const float *__restrict__ in, +template +__global__ void welford_column_first_pass(const InT *__restrict__ in, WelfordData *__restrict__ out, const size_t num_features, const size_t size) { @@ -180,7 +183,7 @@ __global__ void welford_column_first_pass(const float *__restrict__ in, if (feature >= num_features) return; - WelfordData localData; + WelfordData localData{}; const uint32_t base_idx = batch * num_features + feature; const uint32_t output_idx = num_features * blockIdx.y + feature; @@ -195,7 +198,7 @@ __global__ void welford_column_first_pass(const float *__restrict__ in, return; } - float v = in[idx]; + float v = scalar_to_float(in[idx]); update(localData, v, true); } @@ -203,8 +206,9 @@ __global__ void welford_column_first_pass(const float *__restrict__ in, out[output_idx] = localData; } +template __global__ void welford_column_second_pass(const WelfordData *__restrict__ in, - float *__restrict__ out, + OutT *__restrict__ out, const size_t num_features, const uint32_t num_rows, WelfordType type) { @@ -224,21 +228,22 @@ __global__ void welford_column_second_pass(const WelfordData *__restrict__ in, switch (type) { case WelfordType::Mean: { - out[base_idx] = local.mean; + out[base_idx] = scalar_from_float(local.mean); } break; case WelfordType::PopulationVariance: { - out[base_idx] = local.M2 / local.count; + out[base_idx] = scalar_from_float(local.M2 / local.count); } break; case WelfordType::SampleVariance: { - out[base_idx] = local.M2 / (local.count - 1); + out[base_idx] = scalar_from_float(local.M2 / (local.count - 1)); } break; } } -void launch_welford(void *in, void *out, size_t num_features, size_t batch_size, - int32_t dim, WelfordType type) { +void launch_welford(const void *in, DataType in_dtype, void *out, + DataType out_dtype, size_t num_features, + size_t batch_size, int32_t dim, WelfordType type) { if (num_features == 0 || batch_size == 0) return; @@ -257,24 +262,29 @@ void launch_welford(void *in, void *out, size_t num_features, size_t batch_size, dim3 grid_size(feature_blocks, row_chunks); - cudaMalloc(&staging_buffer, - sizeof(WelfordData) * feature_blocks * batch_size); + CHECK_CUDA(cudaMalloc(&staging_buffer, + sizeof(WelfordData) * feature_blocks * batch_size)); - welford_row_first_pass - <<>>(static_cast(in), - static_cast(staging_buffer), - num_features, batch_size); + dispatch_float_dtype(in_dtype, [&]() { + welford_row_first_pass + <<>>( + static_cast(in), + static_cast(staging_buffer), num_features, + batch_size); + }); const int32_t num_iter = (feature_blocks + welford_internal::kBlockDim - 1) / welford_internal::kBlockDim; - welford_row_second_pass - <<>>( - static_cast(staging_buffer), - static_cast(out), static_cast(feature_blocks), - num_iter, type); + dispatch_float_dtype(out_dtype, [&]() { + welford_row_second_pass + <<>>( + static_cast(staging_buffer), + static_cast(out), static_cast(feature_blocks), + num_iter, type); + }); } else if (dim == 1) { const uint32_t col_chunks = (batch_size + welford_internal::kColChunkSize - 1) / @@ -282,20 +292,26 @@ void launch_welford(void *in, void *out, size_t num_features, size_t batch_size, dim3 grid_size(feature_blocks, col_chunks); - cudaMalloc(&staging_buffer, - sizeof(WelfordData) * num_features * col_chunks); - - welford_column_first_pass - <<>>(static_cast(in), - static_cast(staging_buffer), - num_features, num_features * batch_size); - - welford_column_second_pass<<>>( - static_cast(staging_buffer), - static_cast(out), num_features, col_chunks, type); + CHECK_CUDA(cudaMalloc(&staging_buffer, + sizeof(WelfordData) * num_features * col_chunks)); + + dispatch_float_dtype(in_dtype, [&]() { + welford_column_first_pass + <<>>( + static_cast(in), + static_cast(staging_buffer), num_features, + num_features * batch_size); + }); + + dispatch_float_dtype(out_dtype, [&]() { + welford_column_second_pass<<>>( + static_cast(staging_buffer), + static_cast(out), num_features, col_chunks, type); + }); } - cudaFree(staging_buffer); + CHECK_CUDA(cudaGetLastError()); + CHECK_CUDA(cudaFree(staging_buffer)); } } // namespace smollnet From 36b989429a8bb35a3fae80b577e58b7718094220 Mon Sep 17 00:00:00 2001 From: Jacob Domagala Date: Thu, 13 Nov 2025 12:07:00 +0100 Subject: [PATCH 08/13] [#81]: Update benchmark kernel launches --- test/benchmark/mse_benchmark.cpp | 3 ++- test/benchmark/welford_benchmark.cpp | 4 ++-- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/test/benchmark/mse_benchmark.cpp b/test/benchmark/mse_benchmark.cpp index 34ee486..0c5bd79 100644 --- a/test/benchmark/mse_benchmark.cpp +++ b/test/benchmark/mse_benchmark.cpp @@ -143,7 +143,8 @@ BenchmarkResult run_case(const BenchmarkCase &cfg, Tensor loss = zeros({1}, DataType::f32, Device::CUDA); const auto timing = bench::measure_cuda_operation(run_cfg, [&] { - launch_mse(loss.data(), pred.data(), target.data(), cfg.elements); + launch_mse(loss.data(), loss.dtype(), pred.data(), target.data(), + pred.dtype(), cfg.elements); }); const double total_elems = static_cast(cfg.elements); diff --git a/test/benchmark/welford_benchmark.cpp b/test/benchmark/welford_benchmark.cpp index 882e869..9e78e49 100644 --- a/test/benchmark/welford_benchmark.cpp +++ b/test/benchmark/welford_benchmark.cpp @@ -155,8 +155,8 @@ BenchmarkResult run_case(const BenchmarkCase &cfg, Tensor variance = zeros(output_dims, DataType::f32, Device::CUDA); const auto timing = bench::measure_cuda_operation(run_cfg, [&] { - launch_welford(input.data(), variance.data(), cfg.num_features, - cfg.batch_size, mode.dim, + launch_welford(input.data(), input.dtype(), variance.data(), + variance.dtype(), cfg.num_features, cfg.batch_size, mode.dim, WelfordType::PopulationVariance); }); From 11d4c17611f5c69bc5a953faef826395292e6b58 Mon Sep 17 00:00:00 2001 From: Jacob Domagala Date: Thu, 13 Nov 2025 12:08:00 +0100 Subject: [PATCH 09/13] [#81]: Fix CPU sum scalar indexing --- src/tensor.cpp | 13 +++++++------ 1 file changed, 7 insertions(+), 6 deletions(-) diff --git a/src/tensor.cpp b/src/tensor.cpp index de71287..a47ed82 100644 --- a/src/tensor.cpp +++ b/src/tensor.cpp @@ -764,19 +764,20 @@ Tensor sum(const Tensor &t, int64_t dim, bool keep_dim) { if (device == Device::CPU) { for (size_t linear = 0; linear < new_tensor.numel(); ++linear) { - size_t remaining = linear; + int64_t remaining = static_cast(linear); int64_t input_base_offset = 0; int64_t output_offset = 0; - for (int64_t out_dim = s_output.rank - 1; out_dim >= 0; --out_dim) { - const int64_t coord = remaining % s_output.size[out_dim]; - remaining /= s_output.size[out_dim]; + for (int64_t out_dim = s_output.rank; out_dim > 0; --out_dim) { + const int64_t output_axis = out_dim - 1; + const int64_t coord = remaining % s_output.size[output_axis]; + remaining /= s_output.size[output_axis]; - output_offset += coord * s_output.stride[out_dim]; + output_offset += coord * s_output.stride[output_axis]; const bool kept_dim = s_output.rank == s_input.rank; const int64_t input_dim = - kept_dim || out_dim < dim ? out_dim : out_dim + 1; + kept_dim || output_axis < dim ? output_axis : output_axis + 1; input_base_offset += coord * s_input.stride[input_dim]; } From 4679ce1d9a77add6658e3e38b19615152153e9a5 Mon Sep 17 00:00:00 2001 From: Jacob Domagala Date: Thu, 13 Nov 2025 12:09:00 +0100 Subject: [PATCH 10/13] [#81]: Install Python headers for static analysis --- .github/workflows/build_and_sa.yml | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/.github/workflows/build_and_sa.yml b/.github/workflows/build_and_sa.yml index bee6758..82d7ced 100644 --- a/.github/workflows/build_and_sa.yml +++ b/.github/workflows/build_and_sa.yml @@ -67,7 +67,8 @@ jobs: apt-get install -y \ cuda-minimal-build-13-2 \ cuda-cudart-dev-13-2 \ - cuda-cccl-13-2 + cuda-cccl-13-2 \ + python3-dev rm -f cuda-keyring_1.1-1_all.deb EOF From a070cfc776b714b6285ccb0445b0c40a92c34e8e Mon Sep 17 00:00:00 2001 From: Jacob Domagala Date: Thu, 13 Nov 2025 13:23:45 +0100 Subject: [PATCH 11/13] [#81]: Fix SA issue with non-const pointer --- src/tensor.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/tensor.cpp b/src/tensor.cpp index a47ed82..f66660e 100644 --- a/src/tensor.cpp +++ b/src/tensor.cpp @@ -748,7 +748,7 @@ Tensor sum(const Tensor &t, int64_t dim, bool keep_dim) { Tensor new_tensor = zeros(out_dims.data(), new_rank, output_type, device, t.requires_grad()); - auto *srcp = t.data(); + const auto *srcp = t.data(); auto *dst = new_tensor.data(); StrideAndSize s_input{}; From e16aaabe91e371e3e871d79e7aacac81a250c4a7 Mon Sep 17 00:00:00 2001 From: Jacob Domagala Date: Fri, 14 Nov 2025 13:23:45 +0100 Subject: [PATCH 12/13] [#81]: Update Welford benchmark to use fp16 and fp32 --- test/benchmark/welford_benchmark.cpp | 46 ++++++++++++++++++++-------- 1 file changed, 33 insertions(+), 13 deletions(-) diff --git a/test/benchmark/welford_benchmark.cpp b/test/benchmark/welford_benchmark.cpp index 9e78e49..68ff1b7 100644 --- a/test/benchmark/welford_benchmark.cpp +++ b/test/benchmark/welford_benchmark.cpp @@ -18,6 +18,11 @@ struct WelfordBenchmarkMode { int32_t dim; }; +struct WelfordBenchmarkDType { + const char *label; + DataType dtype; +}; + struct WelfordStageEntry { int count; float mean; @@ -76,6 +81,11 @@ constexpr std::array kModes = {{ {"column", 1}, }}; +constexpr std::array kDTypes = {{ + {"fp32", DataType::f32}, + {"fp16", DataType::f16}, +}}; + size_t ceil_div(size_t numerator, size_t denominator) { return (numerator + denominator - 1) / denominator; } @@ -109,10 +119,11 @@ BenchmarkConfig parse_args(int argc, char **argv) { } double bytes_per_iteration(const BenchmarkCase &cfg, - const WelfordBenchmarkMode &mode) { + const WelfordBenchmarkMode &mode, + DataType dtype) { const double input_bytes = static_cast(cfg.batch_size) * static_cast(cfg.num_features) * - sizeof(float); + static_cast(element_size(dtype)); const double stage_entries = mode.dim == 0 ? static_cast(ceil_div(cfg.num_features, @@ -124,7 +135,8 @@ double bytes_per_iteration(const BenchmarkCase &cfg, const double staging_bytes = stage_entries * static_cast(sizeof(WelfordStageEntry)); const double output_bytes = - static_cast(output_elements(cfg, mode)) * sizeof(float); + static_cast(output_elements(cfg, mode)) * + static_cast(element_size(dtype)); // launch_welford is two-pass for both dim=0 and dim=1: // 1) read input, write Welford staging tuples @@ -143,16 +155,17 @@ struct BenchmarkResult { BenchmarkResult run_case(const BenchmarkCase &cfg, const WelfordBenchmarkMode &mode, + const WelfordBenchmarkDType &dtype, const bench::RunConfig &run_cfg) { Tensor input = rand({static_cast(cfg.batch_size), static_cast(cfg.num_features)}, - DataType::f32, Device::CUDA); + dtype.dtype, Device::CUDA); const int64_t output_dims[2] = { mode.dim == 0 ? static_cast(cfg.batch_size) : 1, mode.dim == 0 ? 1 : static_cast(cfg.num_features), }; - Tensor variance = zeros(output_dims, DataType::f32, Device::CUDA); + Tensor variance = zeros(output_dims, dtype.dtype, Device::CUDA); const auto timing = bench::measure_cuda_operation(run_cfg, [&] { launch_welford(input.data(), input.dtype(), variance.data(), @@ -162,7 +175,7 @@ BenchmarkResult run_case(const BenchmarkCase &cfg, const double total_elems = static_cast(cfg.batch_size) * static_cast(cfg.num_features); - const double bytes_per_iter = bytes_per_iteration(cfg, mode); + const double bytes_per_iter = bytes_per_iteration(cfg, mode, dtype.dtype); const double effective_gb_per_sec = (bytes_per_iter / (timing.avg_ms / 1000.0)) / 1.0e9; @@ -176,10 +189,12 @@ BenchmarkResult run_case(const BenchmarkCase &cfg, }; } -void print_case(const WelfordBenchmarkMode &mode, const BenchmarkCase &cfg, +void print_case(const WelfordBenchmarkMode &mode, + const WelfordBenchmarkDType &dtype, const BenchmarkCase &cfg, const BenchmarkResult &result) { bench::print_fields({ bench::field("axis", bench::ansi::kBoldCyan, "{}", mode.label), + bench::field("dtype", bench::ansi::kBoldMagenta, "{}", dtype.label), bench::field("batch", bench::ansi::kBoldBlue, "{:>6}", cfg.batch_size), bench::field("features", bench::ansi::kBoldBlue, "{:>6}", cfg.num_features), @@ -213,17 +228,22 @@ int main(int argc, char **argv) { bench::field("suite_cases", bench::ansi::kBoldGreen, "{}", kDefaultCases.size()), bench::field("axes", bench::ansi::kBoldCyan, "row,column"), + bench::field("dtypes", bench::ansi::kBoldMagenta, "fp32,fp16"), }); for (const auto &bench_case : kDefaultCases) { - for (const auto &mode : kModes) { - const auto result = run_case(bench_case, mode, cfg.run); - print_case(mode, bench_case, result); + for (const auto &dtype : kDTypes) { + for (const auto &mode : kModes) { + const auto result = run_case(bench_case, mode, dtype, cfg.run); + print_case(mode, dtype, bench_case, result); + } } } } else { - for (const auto &mode : kModes) { - const auto result = run_case(cfg.single_case, mode, cfg.run); - print_case(mode, cfg.single_case, result); + for (const auto &dtype : kDTypes) { + for (const auto &mode : kModes) { + const auto result = run_case(cfg.single_case, mode, dtype, cfg.run); + print_case(mode, dtype, cfg.single_case, result); + } } } From 753c8c41f92454fdd4ed5210730fd0d627fa1584 Mon Sep 17 00:00:00 2001 From: Jacob Domagala Date: Fri, 14 Nov 2025 13:23:45 +0100 Subject: [PATCH 13/13] [#81]: Update MSE benchmark to use fp16 and fp32 --- test/benchmark/mse_benchmark.cpp | 42 +++++++++++++++++++++++--------- 1 file changed, 31 insertions(+), 11 deletions(-) diff --git a/test/benchmark/mse_benchmark.cpp b/test/benchmark/mse_benchmark.cpp index 0c5bd79..32f082c 100644 --- a/test/benchmark/mse_benchmark.cpp +++ b/test/benchmark/mse_benchmark.cpp @@ -21,6 +21,11 @@ struct BenchmarkCase { size_t elements; }; +struct BenchmarkDType { + const char *label; + DataType dtype; +}; + struct BenchmarkConfig { BenchmarkCase single_case{1ull << 24}; bench::RunConfig run; @@ -37,6 +42,11 @@ constexpr std::array kDefaultCases = {{ {1ull << 25}, }}; +constexpr std::array kDTypes = {{ + {"fp32", DataType::f32}, + {"fp16", DataType::f16}, +}}; + size_t parse_size_arg(const char *text, const char *name) { char *end = nullptr; const unsigned long long value = std::strtoull(text, &end, 10); @@ -91,8 +101,9 @@ size_t blocks_per_launch(const BenchmarkCase &cfg) { return ceil_div(cfg.elements, static_cast(kMseBlockSize)); } -double input_bytes_per_iteration(const BenchmarkCase &cfg) { - return static_cast(cfg.elements) * 2.0 * sizeof(float); +double input_bytes_per_iteration(const BenchmarkCase &cfg, DataType dtype) { + return static_cast(cfg.elements) * 2.0 * + static_cast(element_size(dtype)); } double flops_per_iteration(const BenchmarkCase &cfg) { @@ -135,12 +146,13 @@ struct BenchmarkResult { }; BenchmarkResult run_case(const BenchmarkCase &cfg, + const BenchmarkDType &dtype, const bench::RunConfig &run_cfg) { - Tensor pred = rand({static_cast(cfg.elements)}, DataType::f32, + Tensor pred = rand({static_cast(cfg.elements)}, dtype.dtype, Device::CUDA); - Tensor target = rand({static_cast(cfg.elements)}, DataType::f32, + Tensor target = rand({static_cast(cfg.elements)}, dtype.dtype, Device::CUDA); - Tensor loss = zeros({1}, DataType::f32, Device::CUDA); + Tensor loss = zeros({1}, dtype.dtype, Device::CUDA); const auto timing = bench::measure_cuda_operation(run_cfg, [&] { launch_mse(loss.data(), loss.dtype(), pred.data(), target.data(), @@ -148,7 +160,8 @@ BenchmarkResult run_case(const BenchmarkCase &cfg, }); const double total_elems = static_cast(cfg.elements); - const double input_bytes_per_iter = input_bytes_per_iteration(cfg); + const double input_bytes_per_iter = + input_bytes_per_iteration(cfg, dtype.dtype); const double flops_per_iter = flops_per_iteration(cfg); const double effective_input_gb_per_sec = (input_bytes_per_iter / (timing.avg_ms / 1000.0)) / 1.0e9; @@ -168,8 +181,10 @@ BenchmarkResult run_case(const BenchmarkCase &cfg, }; } -void print_case(const BenchmarkCase &cfg, const BenchmarkResult &result) { +void print_case(const BenchmarkCase &cfg, const BenchmarkDType &dtype, + const BenchmarkResult &result) { bench::print_fields({ + bench::field("dtype", bench::ansi::kBoldMagenta, "{}", dtype.label), bench::field("elements", bench::ansi::kBoldBlue, "{:>12}", cfg.elements), bench::field("blocks", bench::ansi::kBoldCyan, "{:>9}", @@ -208,14 +223,19 @@ int main(int argc, char **argv) { bench::print_fields({ bench::field("suite_cases", bench::ansi::kBoldGreen, "{}", kDefaultCases.size()), + bench::field("dtypes", bench::ansi::kBoldMagenta, "fp32,fp16"), }); for (const auto &bench_case : kDefaultCases) { - const auto result = run_case(bench_case, cfg.run); - print_case(bench_case, result); + for (const auto &dtype : kDTypes) { + const auto result = run_case(bench_case, dtype, cfg.run); + print_case(bench_case, dtype, result); + } } } else { - const auto result = run_case(cfg.single_case, cfg.run); - print_case(cfg.single_case, result); + for (const auto &dtype : kDTypes) { + const auto result = run_case(cfg.single_case, dtype, cfg.run); + print_case(cfg.single_case, dtype, result); + } } return 0;