Skip to content

Commit 2d3375d

Browse files
authored
Matmul two variables (#33)
* started on matmul * removed blas for now * minor to matmul * removed blas file * clean up blas * minor
1 parent 6adf06f commit 2d3375d

11 files changed

Lines changed: 760 additions & 2 deletions

File tree

‎CMakeLists.txt‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,5 @@
11
cmake_minimum_required(VERSION 3.15)
22
project(DNLP_Diff_Engine C)
3-
43
set(CMAKE_C_STANDARD 99)
54

65
#Set default build type to Release if not specified
@@ -32,6 +31,7 @@ include_directories(${PROJECT_SOURCE_DIR}/include)
3231
# Source files - automatically gather all .c files from src/
3332
file(GLOB_RECURSE SOURCES "src/*.c")
3433

34+
3535
# Create core library
3636
add_library(dnlp_diff ${SOURCES})
3737
target_link_libraries(dnlp_diff m)

‎include/bivariate.h‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,9 @@ expr *new_quad_over_lin(expr *left, expr *right);
1010
expr *new_rel_entr_first_arg_scalar(expr *left, expr *right);
1111
expr *new_rel_entr_second_arg_scalar(expr *left, expr *right);
1212

13+
/* Matrix multiplication: Z = X @ Y */
14+
expr *new_matmul(expr *x, expr *y);
15+
1316
/* Left matrix multiplication: A @ f(x) where A is a constant matrix */
1417
expr *new_left_matmul(expr *u, const CSR_Matrix *A);
1518

‎include/utils/mini_numpy.h‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,4 +20,7 @@ void tile_int(int *result, const int *a, int len, int tiles);
2020
*/
2121
void scaled_ones(double *result, int size, double value);
2222

23+
/* Naive implementation of Z = X @ Y, X is m x k, Y is k x n, Z is m x n */
24+
void mat_mat_mult(const double *X, const double *Y, double *Z, int m, int k, int n);
25+
2326
#endif /* MINI_NUMPY_H */

‎python/atoms/matmul.h‎

Lines changed: 33 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,33 @@
1+
#ifndef ATOM_MATMUL_H
2+
#define ATOM_MATMUL_H
3+
4+
#include "bivariate.h"
5+
#include "common.h"
6+
7+
/* Matrix multiplication: Z = X @ Y */
8+
static PyObject *py_make_matmul(PyObject *self, PyObject *args)
9+
{
10+
PyObject *left_capsule, *right_capsule;
11+
if (!PyArg_ParseTuple(args, "OO", &left_capsule, &right_capsule))
12+
{
13+
return NULL;
14+
}
15+
expr *left = (expr *) PyCapsule_GetPointer(left_capsule, EXPR_CAPSULE_NAME);
16+
expr *right = (expr *) PyCapsule_GetPointer(right_capsule, EXPR_CAPSULE_NAME);
17+
if (!left || !right)
18+
{
19+
PyErr_SetString(PyExc_ValueError, "invalid child capsule");
20+
return NULL;
21+
}
22+
23+
expr *node = new_matmul(left, right);
24+
if (!node)
25+
{
26+
PyErr_SetString(PyExc_RuntimeError, "failed to create matmul node");
27+
return NULL;
28+
}
29+
expr_retain(node); /* Capsule owns a reference */
30+
return PyCapsule_New(node, EXPR_CAPSULE_NAME, expr_capsule_destructor);
31+
}
32+
33+
#endif /* ATOM_MATMUL_H */

‎python/bindings.c‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@
1919
#include "atoms/linear.h"
2020
#include "atoms/log.h"
2121
#include "atoms/logistic.h"
22+
#include "atoms/matmul.h"
2223
#include "atoms/multiply.h"
2324
#include "atoms/neg.h"
2425
#include "atoms/power.h"
@@ -73,6 +74,8 @@ static PyMethodDef DNLPMethods[] = {
7374
{"make_promote", py_make_promote, METH_VARARGS, "Create promote node"},
7475
{"make_multiply", py_make_multiply, METH_VARARGS,
7576
"Create elementwise multiply node"},
77+
{"make_matmul", py_make_matmul, METH_VARARGS,
78+
"Create matrix multiplication node (Z = X @ Y)"},
7679
{"make_const_scalar_mult", py_make_const_scalar_mult, METH_VARARGS,
7780
"Create constant scalar multiplication node (a * f(x))"},
7881
{"make_const_vector_mult", py_make_const_vector_mult, METH_VARARGS,

0 commit comments

Comments
 (0)