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(¤t);
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 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 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}