| 30 | } // anonymous namespace |
| 31 | |
| 32 | void init_contrib_module(py::module& m) { |
| 33 | // KNN |
| 34 | DEF_REQ(KnnRequest); |
| 35 | DEF_RES(KnnResponse); |
| 36 | |
| 37 | m.def("new_knn_request", |
| 38 | [](const std::string& node_type, int32_t k) { |
| 39 | KnnRequest* req = new KnnRequest(node_type, k); |
| 40 | return static_cast<OpRequest*>(req); |
| 41 | }, |
| 42 | py::return_value_policy::reference, |
| 43 | py::arg("node_type"), |
| 44 | py::arg("k")); |
| 45 | |
| 46 | m.def("set_knn_request", |
| 47 | [](OpRequest* req, |
| 48 | int32_t batch_size, |
| 49 | int32_t dimension, |
| 50 | py::object inputs) { |
| 51 | ImportNumpy(); |
| 52 | PyArrayObject* input = reinterpret_cast<PyArrayObject*>(inputs.ptr()); |
| 53 | KnnRequest* knn_req = static_cast<KnnRequest*>(req); |
| 54 | knn_req->Set(reinterpret_cast<float*>(PyArray_DATA(input)), |
| 55 | batch_size, dimension); |
| 56 | }); |
| 57 | |
| 58 | m.def("new_knn_response", |
| 59 | []() { |
| 60 | return static_cast<OpResponse*>(new KnnResponse()); |
| 61 | }, |
| 62 | py::return_value_policy::reference); |
| 63 | |
| 64 | m.def("get_knn_ids", |
| 65 | [](OpResponse* res) { |
| 66 | ImportNumpy(); |
| 67 | KnnResponse* knn_res = static_cast<KnnResponse*>(res); |
| 68 | npy_intp shape[1]; |
| 69 | shape[0] = knn_res->BatchSize() * knn_res->K(); |
| 70 | PyArray_Descr* descr = PyArray_DescrFromType(NPY_INT64); |
| 71 | PyObject* obj = PyArray_Zeros(1, shape, descr, 0); |
| 72 | PyArrayObject* np_array = reinterpret_cast<PyArrayObject*>(obj); |
| 73 | memcpy(PyArray_DATA(np_array), knn_res->Ids(), |
| 74 | shape[0] * INT64_BYTES); |
| 75 | CAST_RETURN(obj); |
| 76 | }, |
| 77 | py::return_value_policy::reference); |
| 78 | |
| 79 | m.def("get_knn_distances", |
| 80 | [](OpResponse* res) { |
| 81 | ImportNumpy(); |
| 82 | KnnResponse* knn_res = static_cast<KnnResponse*>(res); |
| 83 | npy_intp shape[1]; |
| 84 | shape[0] = knn_res->BatchSize() * knn_res->K(); |
| 85 | PyArray_Descr* descr = PyArray_DescrFromType(NPY_FLOAT32); |
| 86 | PyObject* obj = PyArray_Zeros(1, shape, descr, 0); |
| 87 | PyArrayObject* np_array = reinterpret_cast<PyArrayObject*>(obj); |
| 88 | memcpy(PyArray_DATA(np_array), knn_res->Distances(), |
| 89 | shape[0] * FLOAT32_BYTES); |
no test coverage detected