Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
- Notifications
You must be signed in to change notification settings - Fork 51k
Add matrix_multiplication#10045
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Uh oh!
There was an error while loading. Please reload this page.
Add matrix_multiplication #10045
Changes from all commits
7672d1a642af235e6100c38190985a5a84d2cbbe3af233a5419a7879d8040f56c5b418ea541e686dc5a9de92d3d8325b64e6fda58cccf2a95f9331a35f8b7aaa0da39d451d62dabd6c94a118f339b33638656afd552a873641c0632fd6024667055a0d83b212af89bcb2e23081c72a650ce40a234942b3891da5ef52f3b524d265b435efe6e36fa6a84c027395502443718979e30c5585824fcb1bcbd84532e385b1b621984adad844cd487ab314c62cf90902e4151c9765feafbe1dfc8bd1970250211937a5dab909331be05b52439d08715203d8b620ba2cc6c5cfdaa7e9178bfeeb2bf6d95ae42fafcd87b51fac92bef1422b8721f6956f4db7d9cc0b4aa22c19e3d4e7c77057a2b473405288ef62449d8c6f19a486c9306bba009674b109aca9f7b3146446837d087e37dd5a044ddd7fc6965650d8ba3052c27d28a5cb72b27f25131f6278149bFile filter
Filter by extension
Conversations
Uh oh!
There was an error while loading. Please reload this page.
Jump to
Uh oh!
There was an error while loading. Please reload this page.
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,180 @@ | ||
| # @Author : ojas-wani | ||
| # @File : matrix_multiplication_recursion.py | ||
| # @Date : 10/06/2023 | ||
| """ | ||
| Perform matrix multiplication using a recursive algorithm. | ||
| https://en.wikipedia.org/wiki/Matrix_multiplication | ||
| """ | ||
| # type Matrix = list[list[int]] # psf/black currenttly fails on this line | ||
| Matrix = list[list[int]] | ||
cclauss marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. cclauss marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Variable and function names should follow the There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Variable and function names should follow the | ||
| matrix_1_to_4 = [ | ||
| [1, 2], | ||
| [3, 4], | ||
| ] | ||
| matrix_5_to_8 = [ | ||
| [5, 6], | ||
| [7, 8], | ||
| ] | ||
| matrix_5_to_9_high = [ | ||
| [5, 6], | ||
| [7, 8], | ||
| [9], | ||
| ] | ||
| matrix_5_to_9_wide = [ | ||
| [5, 6], | ||
| [7, 8, 9], | ||
| ] | ||
| matrix_count_up = [ | ||
| [1, 2, 3, 4], | ||
| [5, 6, 7, 8], | ||
| [9, 10, 11, 12], | ||
| [13, 14, 15, 16], | ||
| ] | ||
| matrix_unordered = [ | ||
| [5, 8, 1, 2], | ||
| [6, 7, 3, 0], | ||
| [4, 5, 9, 1], | ||
| [2, 6, 10, 14], | ||
| ] | ||
| matrices = ( | ||
| matrix_1_to_4, | ||
| matrix_5_to_8, | ||
| matrix_5_to_9_high, | ||
| matrix_5_to_9_wide, | ||
| matrix_count_up, | ||
| matrix_unordered, | ||
| ) | ||
| def is_square(matrix: Matrix) -> bool: | ||
cclauss marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. | ||
| """ | ||
| >>> is_square([]) | ||
| True | ||
| >>> is_square(matrix_1_to_4) | ||
| True | ||
| >>> is_square(matrix_5_to_9_high) | ||
| False | ||
| """ | ||
| len_matrix = len(matrix) | ||
| return all(len(row) == len_matrix for row in matrix) | ||
| def matrix_multiply(matrix_a: Matrix, matrix_b: Matrix) -> Matrix: | ||
cclauss marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. | ||
| """ | ||
| >>> matrix_multiply(matrix_1_to_4, matrix_5_to_8) | ||
| [[19, 22], [43, 50]] | ||
| """ | ||
| return [ | ||
| [sum(a * b for a, b in zip(row, col)) for col in zip(*matrix_b)] | ||
| for row in matrix_a | ||
| ] | ||
| def matrix_multiply_recursive(matrix_a: Matrix, matrix_b: Matrix) -> Matrix: | ||
| """ | ||
| :param matrix_a: A square Matrix. | ||
| :param matrix_b: Another square Matrix with the same dimensions as matrix_a. | ||
| :return: Result of matrix_a * matrix_b. | ||
| :raises ValueError: If the matrices cannot be multiplied. | ||
| >>> matrix_multiply_recursive([], []) | ||
| [] | ||
| >>> matrix_multiply_recursive(matrix_1_to_4, matrix_5_to_8) | ||
| [[19, 22], [43, 50]] | ||
| >>> matrix_multiply_recursive(matrix_count_up, matrix_unordered) | ||
| [[37, 61, 74, 61], [105, 165, 166, 129], [173, 269, 258, 197], [241, 373, 350, 265]] | ||
| >>> matrix_multiply_recursive(matrix_1_to_4, matrix_5_to_9_wide) | ||
| Traceback (most recent call last): | ||
| ... | ||
| ValueError: Invalid matrix dimensions | ||
| >>> matrix_multiply_recursive(matrix_1_to_4, matrix_5_to_9_high) | ||
| Traceback (most recent call last): | ||
| ... | ||
| ValueError: Invalid matrix dimensions | ||
| >>> matrix_multiply_recursive(matrix_1_to_4, matrix_count_up) | ||
| Traceback (most recent call last): | ||
| ... | ||
| ValueError: Invalid matrix dimensions | ||
| """ | ||
| if not matrix_a or not matrix_b: | ||
| return [] | ||
| if not all( | ||
| (len(matrix_a) == len(matrix_b), is_square(matrix_a), is_square(matrix_b)) | ||
| ): | ||
| raise ValueError("Invalid matrix dimensions") | ||
| # Initialize the result matrix with zeros | ||
| result = [[0] * len(matrix_b[0]) for _ in range(len(matrix_a))] | ||
| # Recursive multiplication of matrices | ||
| def multiply( | ||
cclauss marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. | ||
| i_loop: int, | ||
| j_loop: int, | ||
| k_loop: int, | ||
| matrix_a: Matrix, | ||
| matrix_b: Matrix, | ||
| result: Matrix, | ||
| ) -> None: | ||
| """ | ||
| :param matrix_a: A square Matrix. | ||
| :param matrix_b: Another square Matrix with the same dimensions as matrix_a. | ||
| :param result: Result matrix | ||
| :param i: Index used for iteration during multiplication. | ||
| :param j: Index used for iteration during multiplication. | ||
| :param k: Index used for iteration during multiplication. | ||
| >>> 0 > 1 # Doctests in inner functions are never run | ||
| True | ||
| """ | ||
| if i_loop >= len(matrix_a): | ||
| return | ||
| if j_loop >= len(matrix_b[0]): | ||
| return multiply(i_loop + 1, 0, 0, matrix_a, matrix_b, result) | ||
| if k_loop >= len(matrix_b): | ||
| return multiply(i_loop, j_loop + 1, 0, matrix_a, matrix_b, result) | ||
| result[i_loop][j_loop] += matrix_a[i_loop][k_loop] * matrix_b[k_loop][j_loop] | ||
| return multiply(i_loop, j_loop, k_loop + 1, matrix_a, matrix_b, result) | ||
| # Perform the recursive matrix multiplication | ||
| multiply(0, 0, 0, matrix_a, matrix_b, result) | ||
| return result | ||
| if __name__ == "__main__": | ||
| from doctest import testmod | ||
| failure_count, test_count = testmod() | ||
| if not failure_count: | ||
| matrix_a = matrices[0] | ||
| for matrix_b in matrices[1:]: | ||
| print("Multiplying:") | ||
| for row in matrix_a: | ||
| print(row) | ||
| print("By:") | ||
| for row in matrix_b: | ||
| print(row) | ||
| print("Result:") | ||
| try: | ||
| result = matrix_multiply_recursive(matrix_a, matrix_b) | ||
| for row in result: | ||
| print(row) | ||
| assert result == matrix_multiply(matrix_a, matrix_b) | ||
| except ValueError as e: | ||
| print(f"{e!r}") | ||
| print() | ||
| matrix_a = matrix_b | ||
| print("Benchmark:") | ||
| from functools import partial | ||
| from timeit import timeit | ||
| mytimeit = partial(timeit, globals=globals(), number=100_000) | ||
| for func in ("matrix_multiply", "matrix_multiply_recursive"): | ||
| print(f"{func:>25}(): {mytimeit(f'{func}(matrix_count_up, matrix_unordered)')}") | ||
Uh oh!
There was an error while loading. Please reload this page.