Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 17 additions & 7 deletions src/nns/_nnscore_bindings.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
#include <nanobind/stl/string.h>
#include <nanobind/stl/vector.h>

#include <cstdint>
#include <cstddef>
#include <stdexcept>
#include <string>
Expand All @@ -16,6 +17,7 @@ namespace {

using Vector = nb::ndarray<const double, nb::ndim<1>, nb::c_contig>;
using IntVector = nb::ndarray<const int, nb::ndim<1>, nb::c_contig>;
using Matrix = nb::ndarray<nb::numpy, double, nb::shape<-1, -1>, nb::f_contig>;

std::size_t checked_size(const Vector& x, const char* name) {
const std::size_t n = x.shape(0);
Expand All @@ -25,6 +27,14 @@ std::size_t checked_size(const Vector& x, const char* name) {
return n;
}

Matrix matrix_from_column_major_vector(std::vector<double>&& values, std::size_t dim) {
auto* storage = new std::vector<double>(std::move(values));
nb::capsule owner(storage, [](void* p) noexcept {
delete static_cast<std::vector<double>*>(p);
});
return Matrix(storage->data(), {dim, dim}, owner, {1, static_cast<int64_t>(dim)});
}

void check_same_size(const Vector& x, const Vector& y, const char* x_name, const char* y_name) {
if (checked_size(x, x_name) != checked_size(y, y_name)) {
throw std::invalid_argument(std::string(x_name) + " and " + y_name + " must have the same length.");
Expand Down Expand Up @@ -136,14 +146,14 @@ nb::dict pm_matrix_dict(double degree_lpm,
throw std::invalid_argument("target length must equal d.");
}
checked_flat_matrix_size(variable, n, d, "variable");
const nns::PMMatrixResult result = nns::pm_matrix(degree_lpm, degree_upm, target.data(),
variable.data(), n, d, pop_adj, norm);
nns::PMMatrixResult result = nns::pm_matrix(degree_lpm, degree_upm, target.data(),
variable.data(), n, d, pop_adj, norm);
nb::dict out;
out["cupm"] = result.cupm;
out["dupm"] = result.dupm;
out["dlpm"] = result.dlpm;
out["clpm"] = result.clpm;
out["cov.matrix"] = result.cov;
out["cupm"] = matrix_from_column_major_vector(std::move(result.cupm), result.dim);
out["dupm"] = matrix_from_column_major_vector(std::move(result.dupm), result.dim);
out["dlpm"] = matrix_from_column_major_vector(std::move(result.dlpm), result.dim);
out["clpm"] = matrix_from_column_major_vector(std::move(result.clpm), result.dim);
out["cov.matrix"] = matrix_from_column_major_vector(std::move(result.cov), result.dim);
out["dim"] = result.dim;
return out;
}
Expand Down
27 changes: 12 additions & 15 deletions src/nns/pm_matrix.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,21 +56,11 @@ def pm_matrix(
)
dim = int(native_result["dim"])
result: PMMatrixResult = {
"cupm": np.asarray(native_result["cupm"], dtype=np.float64).reshape(
(dim, dim), order="F"
),
"dupm": np.asarray(native_result["dupm"], dtype=np.float64).reshape(
(dim, dim), order="F"
),
"dlpm": np.asarray(native_result["dlpm"], dtype=np.float64).reshape(
(dim, dim), order="F"
),
"clpm": np.asarray(native_result["clpm"], dtype=np.float64).reshape(
(dim, dim), order="F"
),
"cov.matrix": np.asarray(native_result["cov.matrix"], dtype=np.float64).reshape(
(dim, dim), order="F"
),
"cupm": _native_matrix(native_result["cupm"], dim),
"dupm": _native_matrix(native_result["dupm"], dim),
"dlpm": _native_matrix(native_result["dlpm"], dim),
"clpm": _native_matrix(native_result["clpm"], dim),
"cov.matrix": _native_matrix(native_result["cov.matrix"], dim),
}
if resolved_names is not None:
result["names"] = resolved_names
Expand Down Expand Up @@ -116,6 +106,13 @@ def pm_matrix(
return result


def _native_matrix(value: Any, dim: int) -> NDArray[np.float64]:
matrix = np.asarray(value, dtype=np.float64)
if matrix.shape == (dim, dim):
return matrix
return cast(NDArray[np.float64], matrix.reshape((dim, dim), order="F"))


def _resolve_names(names: Sequence[str] | None, n_cols: int) -> list[str] | None:
if names is None:
return None
Expand Down
4 changes: 4 additions & 0 deletions tests/invariants/test_native_original_src_coverage.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,10 @@ def test_direct_native_partial_moment_smoke(native: ModuleType) -> None:
)
assert native_pm["dim"] == 2
assert set(native_pm) >= {"cupm", "dupm", "dlpm", "clpm", "cov.matrix", "dim"}
for key in ("cupm", "dupm", "dlpm", "clpm", "cov.matrix"):
assert isinstance(native_pm[key], np.ndarray)
assert native_pm[key].shape == (2, 2)
assert native_pm[key].flags.f_contiguous


def test_direct_native_fast_lm_smoke(native: ModuleType) -> None:
Expand Down
Loading