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
2 changes: 2 additions & 0 deletions python/bindings.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -130,6 +130,8 @@ void bind_tensor_creation_functions(pybind11::module &m) {
bind_tensor_creation_overloads(m, "ones", OnesFunctor{});
bind_tensor_creation_overloads(m, "empty", EmptyFunctor{});

m.def("empty_like", &smollnet::empty_like, "tensor"_a,
"requires_grad"_a = false);
m.def("full_like", &smollnet::full_like, "tensor"_a, "value"_a,
"requires_grad"_a = false);
m.def("zeros_like", &smollnet::zeros_like, "tensor"_a,
Expand Down
7 changes: 4 additions & 3 deletions src/autograd.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -287,10 +287,11 @@ TanhFunction::backward(const std::vector<Tensor> &grad_outputs) {
return grad_inputs;
}

// SigmoidFunction implementation
SigmoidFunction::SigmoidFunction(const Tensor &input) {
SigmoidFunction::SigmoidFunction(const Tensor &input,
const Tensor &sigmoid_output) {
inputs = {input};
needs_input_grad = {input.initialized() && input.requires_grad()};
sigmoid_output_data_ = sigmoid_output.data();
}

std::vector<Tensor>
Expand All @@ -304,7 +305,7 @@ 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.dtype(),
sigmoid_output_data_, grad_input.dtype(),
grad_input.numel());
grad_inputs[0] = grad_input;
}
Expand Down
6 changes: 5 additions & 1 deletion src/autograd.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -89,10 +89,14 @@ struct TanhFunction : Function {
};

struct SigmoidFunction : Function {
explicit SigmoidFunction(const Tensor &input);
SigmoidFunction(const Tensor &input, const Tensor &sigmoid_output);
std::vector<Tensor>
backward(const std::vector<Tensor> &grad_outputs) override;
void print() const override { printf("SigmoidFunction\n"); }

private:
// We're not storing Tensor here to avoid refernce cycle
const void *sigmoid_output_data_ = nullptr;
};

struct SumFunction : Function {
Expand Down
11 changes: 8 additions & 3 deletions src/tensor.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -702,8 +702,7 @@ Tensor tanh(const Tensor &t) {
}

Tensor sigmoid(const Tensor &t) {
Tensor new_tensor = empty(t.dims().data(), t.ndims(), t.dtype(), t.device(),
t.requires_grad());
Tensor new_tensor = empty_like(t, t.requires_grad());

if (t.device() == Device::CUDA) {
launch_sigmoid(new_tensor.data(), t.data(), t.dtype(), t.numel());
Expand All @@ -714,7 +713,8 @@ Tensor sigmoid(const Tensor &t) {
1.0f / (1.0f + std::exp(-x)));
}
}
SetupAutograd<SigmoidFunction>(new_tensor, t);
// We reuse the sigmoid redult for efficiency
SetupAutograd<SigmoidFunction>(t, new_tensor, new_tensor);
return new_tensor;
}

Expand Down Expand Up @@ -987,6 +987,11 @@ Tensor rand(const int64_t *dims, size_t rank, DataType t, Device d,
return Tensor{tensor};
}

Tensor empty_like(const Tensor &t, bool requires_grad) {
return empty(t.dims().data(), t.ndims(), t.dtype(), t.device(),
requires_grad);
}

Tensor zeros_like(const Tensor &t, bool requires_grad) {
return zeros(t.dims().data(), t.ndims(), t.dtype(), t.device(),
requires_grad);
Expand Down
2 changes: 2 additions & 0 deletions src/tensor.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -156,6 +156,8 @@ Tensor ones(const int64_t *dims, size_t rank, DataType t, Device d,
bool requires_grad = false);
Tensor rand(const int64_t *dims, size_t rank, DataType t, Device d,
bool requires_grad = false);

Tensor empty_like(const Tensor &t, bool requires_grad = false);
Tensor full_like(const Tensor &t, float value, bool requires_grad = false);
Tensor zeros_like(const Tensor &t, bool requires_grad = false);
Tensor ones_like(const Tensor &t, bool requires_grad = false);
Expand Down
6 changes: 6 additions & 0 deletions utils/tests.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,21 +33,27 @@ def test_tensor_creation():
print(f"2D rand shape: {t2d_rand.dims()}")

print("\nCreating like tensors...")
t2d_empty_like = smollnet.empty_like(t2d_rand)
t2d_full_like = smollnet.full_like(t2d_rand, 3.0)
t2d_zeros_like = smollnet.zeros_like(t2d_rand)
t2d_ones_like = smollnet.ones_like(t2d_rand, requires_grad=True)
t2d_rand_like = smollnet.rand_like(t2d_rand)

assert t2d_empty_like.dims() == t2d_rand.dims()
assert t2d_full_like.dims() == t2d_rand.dims()
assert t2d_zeros_like.dims() == t2d_rand.dims()
assert t2d_ones_like.dims() == t2d_rand.dims()
assert t2d_rand_like.dims() == t2d_rand.dims()
assert t2d_empty_like.dtype() == t2d_rand.dtype()
assert t2d_full_like.dtype() == t2d_rand.dtype()
assert t2d_zeros_like.dtype() == t2d_rand.dtype()
assert t2d_empty_like.device() == t2d_rand.device()
assert t2d_ones_like.device() == t2d_rand.device()
assert not t2d_empty_like.requires_grad()
assert not t2d_zeros_like.requires_grad()
assert t2d_ones_like.requires_grad()

print(f"2D empty_like shape: {t2d_empty_like.dims()}")
print(f"2D full_like shape: {t2d_full_like.dims()}")
print(f"2D zeros_like shape: {t2d_zeros_like.dims()}")
print(f"2D ones_like shape: {t2d_ones_like.dims()}")
Expand Down
Loading