-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathpython_ngsmetal.cpp
More file actions
74 lines (54 loc) · 1.93 KB
/
Copy pathpython_ngsmetal.cpp
File metadata and controls
74 lines (54 loc) · 1.93 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
#include <Foundation/Foundation.hpp>
#include <Metal/Metal.hpp>
#include <comp.hpp>
#include <python_comp.hpp>
#include "ngsmetal.hpp"
#include "metal_vector.hpp"
#include "metal_sparsematrix.hpp"
#include "metal_btdtb.hpp"
using namespace ngsmetal;
namespace ngsmetal {
extern void ComputeBenchmarks();
extern void ComputeBenchmarksLoadCompute();
extern void ComputeBenchmarksLoad();
}
PYBIND11_MODULE(ngsmetal, m)
{
cout << "Loading ngs-metal library" << endl;
ngsmetal::InitNgsMetal();
ngsmetal::InitVectorKernels();
ngsmetal::InitSparseMatrixKernels();
ngsmetal::InitMatMat();
py::class_<MetalVector, BaseVector, shared_ptr<MetalVector>> (m, "MetalVector")
.def(py::init<size_t>(), py::arg("size"))
;
py::class_<MetalSparseMatrix, BaseMatrix, shared_ptr<MetalSparseMatrix>> (m, "MetalSparseMatrix")
.def(py::init<const BaseSparseMatrix&>(), py::arg("mat"))
;
py::class_<MetalBTDTBMatrix, BaseMatrix, shared_ptr<MetalBTDTBMatrix>> (m, "MetalBTDTBMatrix")
.def(py::init<const BaseMatrix&>(), py::arg("mat"))
;
m.def("TestMatMat", [](size_t n, size_t m, size_t k, int runs) {
Matrix<float> a(n,k), b(k,m), c(n,m);
/*
if (n%32 != 0) throw Exception("n must be multiple of 32");
if (m%32 != 0) throw Exception("m must be multiple of 32");
if (k%32 != 0) throw Exception("k must be multiple of 32");
for (int i = 0; i < n; i++)
for (int j = 0; j < k; j++)
a(i,j) = sin(i+j);
for (int i = 0; i < k; i++)
for (int j = 0; j < m; j++)
b(i,j) = sin(i*j);
*/
a = 1; b = 2;
for (int i = 0; i < runs; i++)
Mult (a,b,c);
if (c.Height()*c.Width() < 10000)
cout << "err = " << L2Norm(c-a*b) << endl;
// cout << "c = " << c << endl;
});
m.def("ComputeBenchmarks", &ComputeBenchmarks);
m.def("ComputeBenchmarksLoadCompute", &ComputeBenchmarksLoadCompute);
m.def("ComputeBenchmarksLoad", &ComputeBenchmarksLoad);
}