← Commits · a5cb821f
a5cb821f982fa91f43c49ac1c5af9e7c6e923bd9
diff --git a/test.c b/test.c
index ba64202..487c52e 100644
--- a/test.c
+++ b/test.c
@@ -37,9 +37,30 @@ void test_reshape(void) {
ut_free(a);
}
+void test_elementwise(void) {
+ ut_tensor* a = ut_from_data(1, (int[]){5}, (float[]){1.f, 2.f, 3.f, 4.f, 5.f});
+ ut_tensor* b = ut_from_data(1, (int[]){5}, (float[]){5.f, 4.f, 3.f, 2.f, 1.f});
+
+ ut_tensor* c = ut_add(a, b);
+ for (int i = 0; i < 5; i++) assert(c->data[i] == a->data[i] + b->data[i]);
+
+ ut_tensor* d = ut_mul(a, b);
+ for (int i = 0; i < 5; i++) assert(d->data[i] == a->data[i] * b->data[i]);
+
+ ut_tensor* e = ut_neg(a);
+ for (int i = 0; i < 5; i++) assert(e->data[i] == -a->data[i]);
+
+ ut_free(a);
+ ut_free(b);
+ ut_free(c);
+ ut_free(d);
+ ut_free(e);
+}
+
int main() {
test_shape();
test_lifetime();
test_reshape();
+ test_elementwise();
return 0;
}
diff --git a/utensil.h b/utensil.h
index f893754..aa5172f 100644
--- a/utensil.h
+++ b/utensil.h
@@ -21,6 +21,10 @@ typedef struct ut_tensor {
struct ut_tensor* owner;
} ut_tensor;
+// =========================================================
+// Allocation and lifetime management
+// =========================================================
+
ut_shape ut_shape_new(int ndim, const int* dim) {
ut_shape s = {.ndim = ndim, .nelem = 1};
for (int i = 0; i < ndim; i++) s.shape[i] = dim[i], s.nelem *= dim[i];
@@ -79,4 +83,61 @@ ut_tensor* ut_view(ut_tensor* t, int ndim, const int* dim) {
ut_tensor* ut_retain(ut_tensor* t) { return t->rc++, t; }
+// =========================================================
+// Elementwise operations
+// =========================================================
+static void ew_neg(float* out, const float* a, int n) {
+ for (int i = 0; i < n; i++) out[i] = -a[i];
+}
+static void ew_exp(float* out, const float* a, int n) {
+ for (int i = 0; i < n; i++) out[i] = expf(a[i]);
+}
+static void ew_sigmoid(float* out, const float* a, int n) {
+ for (int i = 0; i < n; i++) out[i] = 1.f / (1.f + expf(-a[i]));
+}
+static void ew_tanh(float* out, const float* a, int n) {
+ for (int i = 0; i < n; i++) out[i] = tanhf(a[i]);
+}
+static void ew_relu(float* out, const float* a, int n) {
+ for (int i = 0; i < n; i++) out[i] = fmaxf(0.f, a[i]);
+}
+static void ew_add(float* out, const float* a, const float* b, int n) {
+ for (int i = 0; i < n; i++) out[i] = a[i] + b[i];
+}
+static void ew_sub(float* out, const float* a, const float* b, int n) {
+ for (int i = 0; i < n; i++) out[i] = a[i] - b[i];
+}
+static void ew_mul(float* out, const float* a, const float* b, int n) {
+ for (int i = 0; i < n; i++) out[i] = a[i] * b[i];
+}
+static void ew_div(float* out, const float* a, const float* b, int n) {
+ for (int i = 0; i < n; i++) out[i] = a[i] / b[i];
+}
+static ut_tensor* ew_unary(ut_tensor* a, void (*fn)(float*, const float*, int)) {
+ ut_tensor* out = ut_alloc(a->shape.ndim, a->shape.shape);
+ fn(out->data, a->data, a->shape.nelem);
+ return out;
+}
+static ut_tensor* ew_binary(ut_tensor* a, ut_tensor* b,
+ void (*fn)(float*, const float*, const float*, int)) {
+ ut_tensor* out = ut_alloc(a->shape.ndim, a->shape.shape);
+ fn(out->data, a->data, b->data, a->shape.nelem);
+ return out;
+}
+
+ut_tensor* ut_neg(ut_tensor* a) { return ew_unary(a, ew_neg); }
+ut_tensor* ut_exp(ut_tensor* a) { return ew_unary(a, ew_exp); }
+ut_tensor* ut_sigmoid(ut_tensor* a) { return ew_unary(a, ew_sigmoid); }
+ut_tensor* ut_tanh(ut_tensor* a) { return ew_unary(a, ew_tanh); }
+ut_tensor* ut_relu(ut_tensor* a) { return ew_unary(a, ew_relu); }
+ut_tensor* ut_add(ut_tensor* a, ut_tensor* b) { return ew_binary(a, b, ew_add); }
+ut_tensor* ut_sub(ut_tensor* a, ut_tensor* b) { return ew_binary(a, b, ew_sub); }
+ut_tensor* ut_mul(ut_tensor* a, ut_tensor* b) { return ew_binary(a, b, ew_mul); }
+ut_tensor* ut_div(ut_tensor* a, ut_tensor* b) { return ew_binary(a, b, ew_div); }
+ut_tensor* ut_scale(ut_tensor* a, float s) {
+ ut_tensor* out = ut_alloc(a->shape.ndim, a->shape.shape);
+ for (int i = 0; i < a->shape.nelem; i++) out->data[i] = a->data[i] * s;
+ return out;
+}
+
#endif // UTENSIL_H