Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion .github/workflows/build_and_sa.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions python/bindings.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -237,6 +237,7 @@ PYBIND11_MODULE(smollnet, m) {
});

py::enum_<smollnet::DataType>(m, "DataType")
.value("f16", smollnet::DataType::f16)
.value("f32", smollnet::DataType::f32)
.export_values();

Expand Down
30 changes: 18 additions & 12 deletions src/autograd.cpp
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
#include "autograd.hpp"
#include "dtype_utils.hpp"
#include "helpers.hpp"
#include "kernels.cuh"

Expand Down Expand Up @@ -116,7 +117,7 @@ SubFunction::backward(const std::vector<Tensor> &grad_outputs) {
}
}

launch_negative(grad.data(), grad.numel());
launch_negative(grad.data(), grad.dtype(), grad.numel());
grad_inputs[1] = grad;
}

Expand Down Expand Up @@ -187,7 +188,7 @@ DivFunction::backward(const std::vector<Tensor> &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;
}

Expand Down Expand Up @@ -239,7 +240,7 @@ ReLUFunction::backward(const std::vector<Tensor> &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;
}
Expand All @@ -257,7 +258,7 @@ GeLUFunction::backward(const std::vector<Tensor> &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;
}
Expand All @@ -278,7 +279,8 @@ TanhFunction::backward(const std::vector<Tensor> &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;
}

Expand All @@ -302,7 +304,8 @@ SigmoidFunction::backward(const std::vector<Tensor> &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;
}

Expand Down Expand Up @@ -341,17 +344,20 @@ MseFunction::MseFunction(const Tensor &pred, const Tensor &tgt) : N(pred.numel()

std::vector<Tensor> MseFunction::backward(const std::vector<Tensor> &grad_outputs) {
ASSERT(grad_outputs.size() == 1, "MSE backward expects 1 grad_output (scalar)");
float c =
*static_cast<float *>(grad_outputs[0].cpu().data()) * (2.f / static_cast<float>(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<float>(N));

std::vector<Tensor> 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;
}
Expand Down Expand Up @@ -392,8 +398,8 @@ LayerNormFunction::backward(const std::vector<Tensor> &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]
Expand Down
103 changes: 103 additions & 0 deletions src/dtype_utils.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,103 @@
#pragma once

#include "helpers.hpp"
#include "types.hpp"

#include <cuda_fp16.h>

#include <cstddef>

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 <typename T> struct ScalarTraits;

template <> struct ScalarTraits<float> {
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 <typename T>
SMOLLNET_HOST_DEVICE float scalar_to_float(T value) {
return ScalarTraits<T>::to_float(value);
}

template <typename T>
SMOLLNET_HOST_DEVICE T scalar_from_float(float value) {
return ScalarTraits<T>::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<const __half *>(data)[index]);
case DataType::f32:
return static_cast<const float *>(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<float *>(data)[index] = value;
return;
default:
ASSERT(false, fmt::format("Unsupported dtype {}", get_name(dtype)));
}

__builtin_unreachable();
}

template <typename Fn> void dispatch_float_dtype(DataType dtype, Fn &&fn) {
switch (dtype) {
case DataType::f16:
fn.template operator()<__half>();
return;
case DataType::f32:
fn.template operator()<float>();
return;
default:
ASSERT(false, fmt::format("Unsupported dtype {}", get_name(dtype)));
}

__builtin_unreachable();
}

} // namespace smollnet
Loading
Loading