commit
2e717ad536
4 changed files with 22 additions and 0 deletions
|
|
@ -1,5 +1,6 @@
|
|||
#include <pybind11/pybind11.h>
|
||||
#include <pybind11/stl.h>
|
||||
#include <pybind11/numpy.h>
|
||||
|
||||
#include <kompute/Kompute.hpp>
|
||||
|
||||
|
|
@ -39,6 +40,20 @@ PYBIND11_MODULE(kp, m) {
|
|||
return std::unique_ptr<kp::Tensor>(new kp::Tensor(data, tensorTypes));
|
||||
}), "Initialiser with list of data components and tensor GPU memory type.")
|
||||
.def("data", &kp::Tensor::data, DOC(kp, Tensor, data))
|
||||
.def("numpy", [](kp::Tensor& self){
|
||||
ssize_t ndim = 1;
|
||||
std::vector<ssize_t> shape = { self.size() };
|
||||
std::vector<ssize_t> strides = { sizeof(float) };
|
||||
|
||||
return py::array(py::buffer_info(
|
||||
self.data().data(),
|
||||
sizeof(float),
|
||||
py::format_descriptor<float>::format(),
|
||||
ndim,
|
||||
shape,
|
||||
strides
|
||||
));
|
||||
}, "Returns stored data as a new numpy array.")
|
||||
.def("__getitem__", [](kp::Tensor &self, size_t index) -> float { return self.data()[index]; },
|
||||
"When only an index is necessary")
|
||||
.def("__setitem__", [](kp::Tensor &self, size_t index, float value) {
|
||||
|
|
|
|||
|
|
@ -1 +1,2 @@
|
|||
pyshader==0.7.0
|
||||
numpy
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import pyshader as ps
|
||||
import kp
|
||||
import numpy as np
|
||||
|
||||
|
||||
def test_array_multiplication():
|
||||
|
|
@ -33,3 +34,4 @@ def test_array_multiplication():
|
|||
mgr.eval_tensor_sync_local_def([tensor_out])
|
||||
|
||||
assert tensor_out.data() == [2.0, 4.0, 6.0]
|
||||
assert np.all(tensor_out.numpy() == [2.0, 4.0, 6.0])
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import os
|
||||
|
||||
import kp
|
||||
import numpy as np
|
||||
|
||||
DIRNAME = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
|
|
@ -22,6 +23,7 @@ def test_opmult():
|
|||
mgr.eval_tensor_sync_local_def([tensor_out])
|
||||
|
||||
assert tensor_out.data() == [2.0, 4.0, 6.0]
|
||||
assert np.all(tensor_out.numpy() == [2.0, 4.0, 6.0])
|
||||
|
||||
def test_opalgobase_data():
|
||||
"""
|
||||
|
|
@ -57,6 +59,7 @@ def test_opalgobase_data():
|
|||
mgr.eval_tensor_sync_local_def([tensor_out])
|
||||
|
||||
assert tensor_out.data() == [2.0, 4.0, 6.0]
|
||||
assert np.all(tensor_out.numpy() == [2.0, 4.0, 6.0])
|
||||
|
||||
|
||||
def test_opalgobase_file():
|
||||
|
|
@ -106,3 +109,4 @@ def test_sequence():
|
|||
seq.eval()
|
||||
|
||||
assert tensor_out.data() == [2.0, 4.0, 6.0]
|
||||
assert np.all(tensor_out.numpy() == [2.0, 4.0, 6.0])
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue