Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- #pragma once
- #include <cassert>
- #include <algorithm>
- #include <memory>
- template<typename T>
- class as_matrix
- {
- size_t nc = 0, nr = 0;
- const bool col_major = true;
- // column major indexer
- auto cm_indexer(size_t i, size_t j) -> T&
- {
- assert(i >= 0 && j >= 0 && i < n_rows && j < n_cols);
- return ptr[j * n_rows + i];
- }
- // row major indexer
- auto rm_indexer(size_t i, size_t j) -> T&
- {
- assert(i >= 0 && j >= 0 && i < n_rows && j < n_cols);
- return ptr[i * n_cols + j];
- }
- public:
- T* ptr = nullptr;
- T& (as_matrix<T>::*indexer) (size_t, size_t) = nullptr;
- const size_t& n_rows = nr, n_cols = nc;
- as_matrix<T>& operator=(const as_matrix<T> matrix)
- {
- nr = matrix.nr;
- nc = matrix.nc;
- ptr = matrix.ptr;
- indexer = (matrix.col_major)? &as_matrix<T>::cm_indexer: &as_matrix<T>::rm_indexer;
- return *this;
- }
- as_matrix(T *ptr, size_t n_rows, size_t n_cols, bool col_major = true)
- : ptr(ptr), col_major(col_major), indexer((col_major)? &as_matrix<T>::cm_indexer: &as_matrix<T>::rm_indexer),
- nr(n_rows), nc(n_cols)
- {
- }
- auto operator()(size_t i, size_t j) const -> const T&
- {
- return ptr[j * n_rows + i];
- }
- auto operator()(size_t i, size_t j) -> T&
- {
- return ptr[j * n_rows + i];
- }
- };
- template<typename T>
- auto copy(const as_matrix<T>& matrix)
- {
- auto owner = std::make_unique<T[]>(matrix.n_rows * matrix.n_cols);
- auto res = as_matrix(owner.get(), matrix.n_rows, matrix.n_cols);
- std::copy(matrix.ptr, matrix.ptr + matrix.n_rows * matrix.n_cols, res.ptr);
- return std::make_tuple(std::move(owner), res);
- }
- template<typename T>
- auto t(as_matrix<T>& matrix)
- {
- auto [owner, temp] = copy(matrix);
- for (size_t i = 0; i < matrix.n_rows; i++)
- {
- for (size_t j = 0; j < matrix.n_cols; j++)
- {
- matrix.ptr[i * matrix.n_cols + j] = temp.ptr[j * matrix.n_rows + i];
- }
- }
- matrix = as_matrix(matrix.ptr, matrix.n_cols, matrix.n_rows);
- }
- template<typename T>
- auto get(as_matrix<T>& matrix, size_t index)
- {
- return (matrix.col_major)? &matrix.ptr[matrix.n_cols * index]: &matrix.ptr[matrix.n_rows * index];
- }
Advertisement
Add Comment
Please, Sign In to add comment