14#include <lagrange/image/Array3D.h>
15#include <lagrange/image/View3D.h>
16#include <lagrange/python/binding.h>
17#include <lagrange/python/tensor_utils.h>
18#include <lagrange/utils/assert.h>
22namespace lagrange::python {
24namespace nb = nanobind;
26using ImageShape = nb::shape<-1, -1, -1>;
28template <
typename Scalar>
29using ImageTensor = nb::ndarray<Scalar, ImageShape, nb::numpy, nb::c_contig, nb::device::cpu>;
34template <
typename Scalar>
35auto tensor_to_image_view(
const ImageTensor<Scalar>& tensor) -> image::experimental::View3D<Scalar>
37 const image::experimental::dextents<size_t, 3> shape{
42 const std::array<size_t, 3> strides{
43 static_cast<size_t>(tensor.stride(1)),
44 static_cast<size_t>(tensor.stride(0)),
45 static_cast<size_t>(tensor.stride(2)),
47 const image::experimental::layout_stride::mapping mapping{shape, strides};
48 image::experimental::View3D<Scalar> view{
49 static_cast<float*
>(tensor.data()),
55template <
typename Scalar>
56void copy_tensor_to_image_view(
57 const ImageTensor<Scalar>& tensor,
58 image::experimental::View3D<Scalar> image)
60 const auto width =
static_cast<unsigned int>(tensor.shape(1));
61 const auto height =
static_cast<unsigned int>(tensor.shape(0));
62 const auto num_channels =
static_cast<unsigned int>(tensor.shape(2));
64 image.extent(0) == width && image.extent(1) == height && image.extent(2) == num_channels,
65 "Tensor and mdspan dimensions do not match");
66 for (
unsigned int j = 0; j < height; j++) {
67 for (
unsigned int i = 0; i < width; i++) {
68 for (
unsigned int c = 0; c < num_channels; c++) {
69 image(i, j, c) = tensor(j, i, c);
75template <
typename Scalar>
76Tensor<Scalar> image_array_to_tensor(image::experimental::Array3D<Scalar>&& image_)
80 using Array = image::experimental::Array3D<Scalar>;
81 auto* image =
new Array(std::move(image_));
82 nb::capsule owner(image, [](
void* p)
noexcept {
delete static_cast<Array*
>(p); });
83 return Tensor<Scalar>(
84 static_cast<Scalar*
>(image->data()),
92 static_cast<int64_t>(image->stride(1)),
93 static_cast<int64_t>(image->stride(0)),
94 static_cast<int64_t>(image->stride(2)),
@ Scalar
Mesh attribute must have exactly 1 channel.
Definition AttributeFwd.h:56
#define la_runtime_assert(...)
Runtime assertion check.
Definition assert.h:177