Merge pull request #103 from alexander-g/numpy

Added numpy() method
This commit is contained in:
Alejandro Saucedo 2020-12-27 13:18:04 +00:00 committed by GitHub
commit 2e717ad536
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
4 changed files with 22 additions and 0 deletions

View file

@ -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) {

View file

@ -1 +1,2 @@
pyshader==0.7.0
numpy

View file

@ -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])

View file

@ -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])