Skip to content

Commit 6fc7a99

Browse files
alisatwat3pre-commit-ci[bot]cclauss
authored
docs: explain Strassen's algorithm complexity in docstrings (#14925)
* docs: explain Strassen's algorithm complexity in docstrings Expand actual_strassen() docstring to describe the divide-and-conquer approach (7 recursive multiplications instead of 8) and note time complexity O(n^log2(7)) ~= O(n^2.807) vs O(n^3) for naive matrix multiplication, plus space complexity O(n^2). Also expand strassen() docstring to explain the padding/trimming wrapper logic. No behavior changes; all existing doctests pass. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Fix Ruff 0.16 lint failures * Fix grammar in docstrings and comments Corrected minor grammatical issues in docstrings and comments. * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: Christian Clauss <cclauss@me.com>
1 parent afce564 commit 6fc7a99

1 file changed

Lines changed: 25 additions & 1 deletion

File tree

‎divide_and_conquer/strassen_matrix_multiplication.py‎

Lines changed: 25 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,7 @@
1+
"""
2+
https://en.wikipedia.org/wiki/Strassen_algorithm
3+
"""
4+
15
from __future__ import annotations
26

37
import math
@@ -32,7 +36,7 @@ def matrix_subtraction(matrix_a: list, matrix_b: list):
3236

3337
def split_matrix(a: list) -> tuple[list, list, list, list]:
3438
"""
35-
Given an even length matrix, returns the top_left, top_right, bot_left, bot_right
39+
Given an even-length matrix, returns the top_left, top_right, bot_left, bot_right
3640
quadrant.
3741
3842
>>> split_matrix([[4,3,2,4],[2,3,1,1],[6,5,4,3],[8,4,1,6]])
@@ -75,6 +79,19 @@ def actual_strassen(matrix_a: list, matrix_b: list) -> list:
7579
"""
7680
Recursive function to calculate the product of two matrices, using the Strassen
7781
Algorithm. It only supports square matrices of any size that is a power of 2.
82+
83+
Strassen's algorithm reduces the number of recursive multiplications needed to
84+
multiply two n x n matrices from the 8 required by the naive divide-and-conquer
85+
approach down to 7, at the cost of a few extra matrix additions/subtractions
86+
(which are cheaper, O(n^2), operations). Each matrix is split into four
87+
(n/2) x (n/2) quadrants; 7 products of quadrant combinations are computed
88+
recursively, and those products are combined with additions/subtractions to
89+
form the four quadrants of the result.
90+
91+
Time complexity: O(n^log2(7)) ~= O(n^2.807), an improvement over the O(n^3) of
92+
the standard/naive matrix multiplication algorithm.
93+
Space complexity: O(n^2) for storing the intermediate quadrant matrices, plus
94+
O(log n) recursion stack depth.
7895
"""
7996
if matrix_dimensions(matrix_a) == (2, 2):
8097
return default_matrix_multiplication(matrix_a, matrix_b)
@@ -106,6 +123,13 @@ def actual_strassen(matrix_a: list, matrix_b: list) -> list:
106123

107124
def strassen(matrix1: list, matrix2: list) -> list:
108125
"""
126+
Multiplies two matrices using Strassen's algorithm, which runs in
127+
O(n^log2(7)) ~= O(n^2.807) time, compared to O(n^3) for naive matrix
128+
multiplication. This implementation pads both input matrices with zeros
129+
until they are square matrices whose dimension is a power of 2 (required
130+
by the divide-and-conquer recursion in actual_strassen), performs the
131+
multiplication, then trims the padding back off the result.
132+
109133
>>> strassen([[2,1,3],[3,4,6],[1,4,2],[7,6,7]], [[4,2,3,4],[2,1,1,1],[8,6,4,2]])
110134
[[34, 23, 19, 15], [68, 46, 37, 28], [28, 18, 15, 12], [96, 62, 55, 48]]
111135
>>> strassen([[3,7,5,6,9],[1,5,3,7,8],[1,4,4,5,7]], [[2,4],[5,2],[1,7],[5,5],[7,8]])

0 commit comments

Comments
 (0)