Coverage for /usr/lib/python3/dist-packages/sympy/matrices/expressions/kronecker.py: 28%
163 statements
« prev ^ index » next coverage.py v7.9.1, created at 2025-06-14 15:55 +0200
« prev ^ index » next coverage.py v7.9.1, created at 2025-06-14 15:55 +0200
1"""Implementation of the Kronecker product"""
2from functools import reduce
3from math import prod
5from sympy.core import Mul, sympify
6from sympy.functions import adjoint
7from sympy.matrices.common import ShapeError
8from sympy.matrices.expressions.matexpr import MatrixExpr
9from sympy.matrices.expressions.transpose import transpose
10from sympy.matrices.expressions.special import Identity
11from sympy.matrices.matrices import MatrixBase
12from sympy.strategies import (
13 canon, condition, distribute, do_one, exhaust, flatten, typed, unpack)
14from sympy.strategies.traverse import bottom_up
15from sympy.utilities import sift
17from .matadd import MatAdd
18from .matmul import MatMul
19from .matpow import MatPow
22def kronecker_product(*matrices):
23 """
24 The Kronecker product of two or more arguments.
26 This computes the explicit Kronecker product for subclasses of
27 ``MatrixBase`` i.e. explicit matrices. Otherwise, a symbolic
28 ``KroneckerProduct`` object is returned.
31 Examples
32 ========
34 For ``MatrixSymbol`` arguments a ``KroneckerProduct`` object is returned.
35 Elements of this matrix can be obtained by indexing, or for MatrixSymbols
36 with known dimension the explicit matrix can be obtained with
37 ``.as_explicit()``
39 >>> from sympy import kronecker_product, MatrixSymbol
40 >>> A = MatrixSymbol('A', 2, 2)
41 >>> B = MatrixSymbol('B', 2, 2)
42 >>> kronecker_product(A)
43 A
44 >>> kronecker_product(A, B)
45 KroneckerProduct(A, B)
46 >>> kronecker_product(A, B)[0, 1]
47 A[0, 0]*B[0, 1]
48 >>> kronecker_product(A, B).as_explicit()
49 Matrix([
50 [A[0, 0]*B[0, 0], A[0, 0]*B[0, 1], A[0, 1]*B[0, 0], A[0, 1]*B[0, 1]],
51 [A[0, 0]*B[1, 0], A[0, 0]*B[1, 1], A[0, 1]*B[1, 0], A[0, 1]*B[1, 1]],
52 [A[1, 0]*B[0, 0], A[1, 0]*B[0, 1], A[1, 1]*B[0, 0], A[1, 1]*B[0, 1]],
53 [A[1, 0]*B[1, 0], A[1, 0]*B[1, 1], A[1, 1]*B[1, 0], A[1, 1]*B[1, 1]]])
55 For explicit matrices the Kronecker product is returned as a Matrix
57 >>> from sympy import Matrix, kronecker_product
58 >>> sigma_x = Matrix([
59 ... [0, 1],
60 ... [1, 0]])
61 ...
62 >>> Isigma_y = Matrix([
63 ... [0, 1],
64 ... [-1, 0]])
65 ...
66 >>> kronecker_product(sigma_x, Isigma_y)
67 Matrix([
68 [ 0, 0, 0, 1],
69 [ 0, 0, -1, 0],
70 [ 0, 1, 0, 0],
71 [-1, 0, 0, 0]])
73 See Also
74 ========
75 KroneckerProduct
77 """
78 if not matrices:
79 raise TypeError("Empty Kronecker product is undefined")
80 if len(matrices) == 1:
81 return matrices[0]
82 else:
83 return KroneckerProduct(*matrices).doit()
86class KroneckerProduct(MatrixExpr):
87 """
88 The Kronecker product of two or more arguments.
90 The Kronecker product is a non-commutative product of matrices.
91 Given two matrices of dimension (m, n) and (s, t) it produces a matrix
92 of dimension (m s, n t).
94 This is a symbolic object that simply stores its argument without
95 evaluating it. To actually compute the product, use the function
96 ``kronecker_product()`` or call the ``.doit()`` or ``.as_explicit()``
97 methods.
99 >>> from sympy import KroneckerProduct, MatrixSymbol
100 >>> A = MatrixSymbol('A', 5, 5)
101 >>> B = MatrixSymbol('B', 5, 5)
102 >>> isinstance(KroneckerProduct(A, B), KroneckerProduct)
103 True
104 """
105 is_KroneckerProduct = True
107 def __new__(cls, *args, check=True):
108 args = list(map(sympify, args))
109 if all(a.is_Identity for a in args):
110 ret = Identity(prod(a.rows for a in args))
111 if all(isinstance(a, MatrixBase) for a in args):
112 return ret.as_explicit()
113 else:
114 return ret
116 if check:
117 validate(*args)
118 return super().__new__(cls, *args)
120 @property
121 def shape(self):
122 rows, cols = self.args[0].shape
123 for mat in self.args[1:]:
124 rows *= mat.rows
125 cols *= mat.cols
126 return (rows, cols)
128 def _entry(self, i, j, **kwargs):
129 result = 1
130 for mat in reversed(self.args):
131 i, m = divmod(i, mat.rows)
132 j, n = divmod(j, mat.cols)
133 result *= mat[m, n]
134 return result
136 def _eval_adjoint(self):
137 return KroneckerProduct(*list(map(adjoint, self.args))).doit()
139 def _eval_conjugate(self):
140 return KroneckerProduct(*[a.conjugate() for a in self.args]).doit()
142 def _eval_transpose(self):
143 return KroneckerProduct(*list(map(transpose, self.args))).doit()
145 def _eval_trace(self):
146 from .trace import trace
147 return Mul(*[trace(a) for a in self.args])
149 def _eval_determinant(self):
150 from .determinant import det, Determinant
151 if not all(a.is_square for a in self.args):
152 return Determinant(self)
154 m = self.rows
155 return Mul(*[det(a)**(m/a.rows) for a in self.args])
157 def _eval_inverse(self):
158 try:
159 return KroneckerProduct(*[a.inverse() for a in self.args])
160 except ShapeError:
161 from sympy.matrices.expressions.inverse import Inverse
162 return Inverse(self)
164 def structurally_equal(self, other):
165 '''Determine whether two matrices have the same Kronecker product structure
167 Examples
168 ========
170 >>> from sympy import KroneckerProduct, MatrixSymbol, symbols
171 >>> m, n = symbols(r'm, n', integer=True)
172 >>> A = MatrixSymbol('A', m, m)
173 >>> B = MatrixSymbol('B', n, n)
174 >>> C = MatrixSymbol('C', m, m)
175 >>> D = MatrixSymbol('D', n, n)
176 >>> KroneckerProduct(A, B).structurally_equal(KroneckerProduct(C, D))
177 True
178 >>> KroneckerProduct(A, B).structurally_equal(KroneckerProduct(D, C))
179 False
180 >>> KroneckerProduct(A, B).structurally_equal(C)
181 False
182 '''
183 # Inspired by BlockMatrix
184 return (isinstance(other, KroneckerProduct)
185 and self.shape == other.shape
186 and len(self.args) == len(other.args)
187 and all(a.shape == b.shape for (a, b) in zip(self.args, other.args)))
189 def has_matching_shape(self, other):
190 '''Determine whether two matrices have the appropriate structure to bring matrix
191 multiplication inside the KroneckerProdut
193 Examples
194 ========
195 >>> from sympy import KroneckerProduct, MatrixSymbol, symbols
196 >>> m, n = symbols(r'm, n', integer=True)
197 >>> A = MatrixSymbol('A', m, n)
198 >>> B = MatrixSymbol('B', n, m)
199 >>> KroneckerProduct(A, B).has_matching_shape(KroneckerProduct(B, A))
200 True
201 >>> KroneckerProduct(A, B).has_matching_shape(KroneckerProduct(A, B))
202 False
203 >>> KroneckerProduct(A, B).has_matching_shape(A)
204 False
205 '''
206 return (isinstance(other, KroneckerProduct)
207 and self.cols == other.rows
208 and len(self.args) == len(other.args)
209 and all(a.cols == b.rows for (a, b) in zip(self.args, other.args)))
211 def _eval_expand_kroneckerproduct(self, **hints):
212 return flatten(canon(typed({KroneckerProduct: distribute(KroneckerProduct, MatAdd)}))(self))
214 def _kronecker_add(self, other):
215 if self.structurally_equal(other):
216 return self.__class__(*[a + b for (a, b) in zip(self.args, other.args)])
217 else:
218 return self + other
220 def _kronecker_mul(self, other):
221 if self.has_matching_shape(other):
222 return self.__class__(*[a*b for (a, b) in zip(self.args, other.args)])
223 else:
224 return self * other
226 def doit(self, **hints):
227 deep = hints.get('deep', True)
228 if deep:
229 args = [arg.doit(**hints) for arg in self.args]
230 else:
231 args = self.args
232 return canonicalize(KroneckerProduct(*args))
235def validate(*args):
236 if not all(arg.is_Matrix for arg in args):
237 raise TypeError("Mix of Matrix and Scalar symbols")
240# rules
242def extract_commutative(kron):
243 c_part = []
244 nc_part = []
245 for arg in kron.args:
246 c, nc = arg.args_cnc()
247 c_part.extend(c)
248 nc_part.append(Mul._from_args(nc))
250 c_part = Mul(*c_part)
251 if c_part != 1:
252 return c_part*KroneckerProduct(*nc_part)
253 return kron
256def matrix_kronecker_product(*matrices):
257 """Compute the Kronecker product of a sequence of SymPy Matrices.
259 This is the standard Kronecker product of matrices [1].
261 Parameters
262 ==========
264 matrices : tuple of MatrixBase instances
265 The matrices to take the Kronecker product of.
267 Returns
268 =======
270 matrix : MatrixBase
271 The Kronecker product matrix.
273 Examples
274 ========
276 >>> from sympy import Matrix
277 >>> from sympy.matrices.expressions.kronecker import (
278 ... matrix_kronecker_product)
280 >>> m1 = Matrix([[1,2],[3,4]])
281 >>> m2 = Matrix([[1,0],[0,1]])
282 >>> matrix_kronecker_product(m1, m2)
283 Matrix([
284 [1, 0, 2, 0],
285 [0, 1, 0, 2],
286 [3, 0, 4, 0],
287 [0, 3, 0, 4]])
288 >>> matrix_kronecker_product(m2, m1)
289 Matrix([
290 [1, 2, 0, 0],
291 [3, 4, 0, 0],
292 [0, 0, 1, 2],
293 [0, 0, 3, 4]])
295 References
296 ==========
298 .. [1] https://en.wikipedia.org/wiki/Kronecker_product
299 """
300 # Make sure we have a sequence of Matrices
301 if not all(isinstance(m, MatrixBase) for m in matrices):
302 raise TypeError(
303 'Sequence of Matrices expected, got: %s' % repr(matrices)
304 )
306 # Pull out the first element in the product.
307 matrix_expansion = matrices[-1]
308 # Do the kronecker product working from right to left.
309 for mat in reversed(matrices[:-1]):
310 rows = mat.rows
311 cols = mat.cols
312 # Go through each row appending kronecker product to.
313 # running matrix_expansion.
314 for i in range(rows):
315 start = matrix_expansion*mat[i*cols]
316 # Go through each column joining each item
317 for j in range(cols - 1):
318 start = start.row_join(
319 matrix_expansion*mat[i*cols + j + 1]
320 )
321 # If this is the first element, make it the start of the
322 # new row.
323 if i == 0:
324 next = start
325 else:
326 next = next.col_join(start)
327 matrix_expansion = next
329 MatrixClass = max(matrices, key=lambda M: M._class_priority).__class__
330 if isinstance(matrix_expansion, MatrixClass):
331 return matrix_expansion
332 else:
333 return MatrixClass(matrix_expansion)
336def explicit_kronecker_product(kron):
337 # Make sure we have a sequence of Matrices
338 if not all(isinstance(m, MatrixBase) for m in kron.args):
339 return kron
341 return matrix_kronecker_product(*kron.args)
344rules = (unpack,
345 explicit_kronecker_product,
346 flatten,
347 extract_commutative)
349canonicalize = exhaust(condition(lambda x: isinstance(x, KroneckerProduct),
350 do_one(*rules)))
353def _kronecker_dims_key(expr):
354 if isinstance(expr, KroneckerProduct):
355 return tuple(a.shape for a in expr.args)
356 else:
357 return (0,)
360def kronecker_mat_add(expr):
361 args = sift(expr.args, _kronecker_dims_key)
362 nonkrons = args.pop((0,), None)
363 if not args:
364 return expr
366 krons = [reduce(lambda x, y: x._kronecker_add(y), group)
367 for group in args.values()]
369 if not nonkrons:
370 return MatAdd(*krons)
371 else:
372 return MatAdd(*krons) + nonkrons
375def kronecker_mat_mul(expr):
376 # modified from block matrix code
377 factor, matrices = expr.as_coeff_matrices()
379 i = 0
380 while i < len(matrices) - 1:
381 A, B = matrices[i:i+2]
382 if isinstance(A, KroneckerProduct) and isinstance(B, KroneckerProduct):
383 matrices[i] = A._kronecker_mul(B)
384 matrices.pop(i+1)
385 else:
386 i += 1
388 return factor*MatMul(*matrices)
391def kronecker_mat_pow(expr):
392 if isinstance(expr.base, KroneckerProduct) and all(a.is_square for a in expr.base.args):
393 return KroneckerProduct(*[MatPow(a, expr.exp) for a in expr.base.args])
394 else:
395 return expr
398def combine_kronecker(expr):
399 """Combine KronekeckerProduct with expression.
401 If possible write operations on KroneckerProducts of compatible shapes
402 as a single KroneckerProduct.
404 Examples
405 ========
407 >>> from sympy.matrices.expressions import combine_kronecker
408 >>> from sympy import MatrixSymbol, KroneckerProduct, symbols
409 >>> m, n = symbols(r'm, n', integer=True)
410 >>> A = MatrixSymbol('A', m, n)
411 >>> B = MatrixSymbol('B', n, m)
412 >>> combine_kronecker(KroneckerProduct(A, B)*KroneckerProduct(B, A))
413 KroneckerProduct(A*B, B*A)
414 >>> combine_kronecker(KroneckerProduct(A, B)+KroneckerProduct(B.T, A.T))
415 KroneckerProduct(A + B.T, B + A.T)
416 >>> C = MatrixSymbol('C', n, n)
417 >>> D = MatrixSymbol('D', m, m)
418 >>> combine_kronecker(KroneckerProduct(C, D)**m)
419 KroneckerProduct(C**m, D**m)
420 """
421 def haskron(expr):
422 return isinstance(expr, MatrixExpr) and expr.has(KroneckerProduct)
424 rule = exhaust(
425 bottom_up(exhaust(condition(haskron, typed(
426 {MatAdd: kronecker_mat_add,
427 MatMul: kronecker_mat_mul,
428 MatPow: kronecker_mat_pow})))))
429 result = rule(expr)
430 doit = getattr(result, 'doit', None)
431 if doit is not None:
432 return doit()
433 else:
434 return result