Skip to main content

feanor_math/algorithms/linsolve/
extension.rs

1use std::alloc::Allocator;
2
3use super::{LinSolveRing, SolveResult};
4use crate::matrix::*;
5use crate::ring::*;
6use crate::rings::extension::{FreeAlgebra, FreeAlgebraStore};
7use crate::seq::*;
8
9#[stability::unstable(feature = "enable")]
10pub fn solve_right_over_extension<R, V1, V2, V3, A>(
11    ring: R,
12    lhs: SubmatrixMut<V1, El<R>>,
13    rhs: SubmatrixMut<V2, El<R>>,
14    mut out: SubmatrixMut<V3, El<R>>,
15    allocator: A,
16) -> SolveResult
17where
18    R: RingStore,
19    R::Type: FreeAlgebra,
20    <<R::Type as RingExtension>::BaseRing as RingStore>::Type: LinSolveRing,
21    V1: AsPointerToSlice<El<R>>,
22    V2: AsPointerToSlice<El<R>>,
23    V3: AsPointerToSlice<El<R>>,
24    A: Allocator,
25{
26    assert_eq!(lhs.row_count(), rhs.row_count());
27    assert_eq!(lhs.col_count(), out.row_count());
28    assert_eq!(rhs.col_count(), out.col_count());
29
30    let mut expanded_lhs = OwnedMatrix::zero_in(
31        lhs.row_count() * ring.rank(),
32        lhs.col_count() * ring.rank(),
33        ring.base_ring(),
34        &allocator,
35    );
36    let mut current;
37    let g = ring.canonical_gen();
38    for i in 0..lhs.row_count() {
39        for j in 0..lhs.col_count() {
40            current = ring.clone_el(lhs.at(i, j));
41            for l in 0..ring.rank() {
42                let current_wrt_basis = ring.wrt_canonical_basis(&current);
43                for k in 0..ring.rank() {
44                    *expanded_lhs.at_mut(i * ring.rank() + k, j * ring.rank() + l) = current_wrt_basis.at(k);
45                }
46                drop(current_wrt_basis);
47                ring.mul_assign_ref(&mut current, &g);
48            }
49        }
50    }
51
52    let mut expanded_rhs = OwnedMatrix::zero_in(
53        rhs.row_count() * ring.rank(),
54        rhs.col_count(),
55        ring.base_ring(),
56        &allocator,
57    );
58    for i in 0..rhs.row_count() {
59        for j in 0..rhs.col_count() {
60            let value_wrt_basis = ring.wrt_canonical_basis(rhs.at(i, j));
61            for k in 0..ring.rank() {
62                *expanded_rhs.at_mut(i * ring.rank() + k, j) = value_wrt_basis.at(k);
63            }
64        }
65    }
66
67    let mut solution = OwnedMatrix::zero_in(
68        lhs.col_count() * ring.rank(),
69        rhs.col_count(),
70        ring.base_ring(),
71        &allocator,
72    );
73    let sol = ring.base_ring().get_ring().solve_right(
74        expanded_lhs.data_mut(),
75        expanded_rhs.data_mut(),
76        solution.data_mut(),
77        &allocator,
78    );
79
80    if !sol.is_solved() {
81        return sol;
82    }
83
84    for i in 0..lhs.col_count() {
85        for j in 0..rhs.col_count() {
86            let res_value = ring.from_canonical_basis(
87                (0..ring.rank()).map(|k| ring.base_ring().clone_el(solution.at(i * ring.rank() + k, j))),
88            );
89            *out.at_mut(i, j) = res_value;
90        }
91    }
92
93    return sol;
94}
95
96#[cfg(test)]
97use std::alloc::Global;
98
99#[cfg(test)]
100use crate::algorithms::matmul::{MatmulAlgorithm, STANDARD_MATMUL};
101#[cfg(test)]
102use crate::assert_matrix_eq;
103#[cfg(test)]
104use crate::rings::extension::extension_impl::FreeAlgebraImpl;
105#[cfg(test)]
106use crate::rings::zn::zn_static;
107
108#[test]
109fn test_solve() {
110    let base_ring = zn_static::Zn::<15>::RING;
111    // Z_15[X]/(X^3 + X^2 + 1);  X^3 + X^2 + 1 = (X + 2)(X + 2X + 2) mod 3, but it is irreducible
112    // mod 5
113    let ring = FreeAlgebraImpl::new(base_ring, 3, vec![14, 0, 14]);
114    let el = |x0: u64, x1: u64, x2: u64| ring.from_canonical_basis([x0, x1, x2]);
115
116    let data_A = [
117        vec![el(1, 0, 0), el(0, 0, 0)],
118        vec![el(2, 1, 0), el(0, 0, 0)],
119        vec![el(0, 0, 0), el(0, 1, 0)],
120    ];
121    let data_B = [vec![el(10, 10, 5)], vec![el(0, 0, 0)], vec![el(1, 0, 0)]];
122    let mut A = OwnedMatrix::from_fn_in(3, 2, |i, j| ring.clone_el(&data_A[i][j]), Global);
123    let mut B = OwnedMatrix::from_fn_in(3, 1, |i, j| ring.clone_el(&data_B[i][j]), Global);
124    let mut sol: OwnedMatrix<_> = OwnedMatrix::zero(2, 1, &ring);
125
126    solve_right_over_extension(&ring, A.data_mut(), B.data_mut(), sol.data_mut(), Global).assert_solved();
127
128    let A = OwnedMatrix::from_fn_in(3, 2, |i, j| ring.clone_el(&data_A[i][j]), Global);
129    let B = OwnedMatrix::from_fn_in(3, 1, |i, j| ring.clone_el(&data_B[i][j]), Global);
130    let mut prod: OwnedMatrix<_> = OwnedMatrix::zero(3, 1, &ring);
131    STANDARD_MATMUL.matmul(
132        TransposableSubmatrix::from(A.data()),
133        TransposableSubmatrix::from(sol.data()),
134        TransposableSubmatrixMut::from(prod.data_mut()),
135        &ring,
136    );
137
138    assert_matrix_eq!(&ring, &B, &prod);
139
140    let data_B = [vec![el(8, 8, 3)], vec![el(0, 0, 0)], vec![el(1, 0, 0)]];
141    let mut A = OwnedMatrix::from_fn_in(3, 2, |i, j| ring.clone_el(&data_A[i][j]), Global);
142    let mut B = OwnedMatrix::from_fn_in(3, 1, |i, j| ring.clone_el(&data_B[i][j]), Global);
143    let mut sol: OwnedMatrix<_> = OwnedMatrix::zero(2, 1, &ring);
144    assert!(!solve_right_over_extension(&ring, A.data_mut(), B.data_mut(), sol.data_mut(), Global).is_solved());
145}
146
147#[test]
148fn test_invert() {
149    let base_ring = zn_static::Zn::<15>::RING;
150    // Z_15[X]/(X^3 + X^2 + 1);  X^3 + X^2 + 1 = (X + 2)(X + 2X + 2) mod 3, but it is irreducible
151    // mod 5
152    let ring = FreeAlgebraImpl::new(base_ring, 3, [14, 0, 14]);
153
154    let matrix = OwnedMatrix::from_fn(2, 2, |i, j| {
155        if i == 0 || j == 0 {
156            ring.one()
157        } else {
158            ring.sub(ring.canonical_gen(), ring.one())
159        }
160    });
161    let mut inverse = OwnedMatrix::zero(2, 2, &ring);
162    solve_right_over_extension(
163        &ring,
164        matrix.clone_matrix(&ring).data_mut(),
165        OwnedMatrix::identity(2, 2, &ring).data_mut(),
166        inverse.data_mut(),
167        Global,
168    )
169    .assert_solved();
170
171    let mut result = OwnedMatrix::zero(2, 2, &ring);
172    STANDARD_MATMUL.matmul(
173        TransposableSubmatrix::from(matrix.data()),
174        TransposableSubmatrix::from(inverse.data()),
175        TransposableSubmatrixMut::from(result.data_mut()),
176        &ring,
177    );
178
179    assert_matrix_eq!(&ring, OwnedMatrix::identity(2, 2, &ring), result);
180}