116 lines
3.6 KiB
Text
116 lines
3.6 KiB
Text
(* ::Package:: *)
|
||
|
||
(* --------------------------------------------------------------- *)
|
||
(* Strassen matrix multiplication – Mathematica version *)
|
||
(* --------------------------------------------------------------- *)
|
||
|
||
ClearAll[Strassen]
|
||
|
||
Strassen[A_?MatrixQ, B_?MatrixQ] :=
|
||
Module[{n = Length[A],
|
||
a11, a12, a21, a22,
|
||
b11, b12, b21, b22,
|
||
p1, p2, p3, p4, p5, p6, p7,
|
||
c11, c12, c21, c22},
|
||
|
||
(* Base case – 1×1 matrix (scalar) *)
|
||
If[n == 1,
|
||
Return[{{A[[1, 1]]*B[[1, 1]]}}];
|
||
];
|
||
|
||
(* ---------------------------------------------------------------- *)
|
||
(* 2‑level block‐view of the matrices (same as Julia's @views slicing) *)
|
||
(* ---------------------------------------------------------------- *)
|
||
a11 = A[[;; n/2, ;; n/2]];
|
||
a12 = A[[;; n/2, n/2 + 1 ;;]];
|
||
a21 = A[[n/2 + 1 ;;, ;; n/2]];
|
||
a22 = A[[n/2 + 1 ;;, n/2 + 1 ;;]];
|
||
|
||
b11 = B[[;; n/2, ;; n/2]];
|
||
b12 = B[[;; n/2, n/2 + 1 ;;]];
|
||
b21 = B[[n/2 + 1 ;;, ;; n/2]];
|
||
b22 = B[[n/2 + 1 ;;, n/2 + 1 ;;]];
|
||
|
||
(* --------------------------------------------------------------- *)
|
||
(* 7 recursive Strassen products *)
|
||
(* --------------------------------------------------------------- *)
|
||
p1 = Strassen[a12 - a22, b21 + b22];
|
||
p2 = Strassen[a11 + a22, b11 + b22];
|
||
p3 = Strassen[a11 - a21, b11 + b12];
|
||
p4 = Strassen[a11 + a12, b22];
|
||
p5 = Strassen[a11, b12 - b22];
|
||
p6 = Strassen[a22, b21 - b11];
|
||
p7 = Strassen[a21 + a22, b11];
|
||
|
||
(* --------------------------------------------------------------- *)
|
||
(* Assemble the four quadrants of the result *)
|
||
(* --------------------------------------------------------------- *)
|
||
c11 = p1 + p2 - p4 + p6;
|
||
c12 = p4 + p5;
|
||
c21 = p6 + p7;
|
||
c22 = p2 - p3 + p5 - p7;
|
||
|
||
(* Join the four blocks back together *)
|
||
Join[
|
||
Join[c11, c12, 2],
|
||
Join[c21, c22, 2],
|
||
1
|
||
]
|
||
];
|
||
|
||
|
||
(* --------------------------------------------------------------- *)
|
||
(* Helper that mimics the Julia `intprint` *)
|
||
(* --------------------------------------------------------------- *)
|
||
ClearAll[intPrint]
|
||
|
||
intPrint[title_String, mat_?MatrixQ] :=
|
||
Module[{rounded},
|
||
(* round each entry to 8 decimal digits, then coerce to an integer
|
||
when the rounded value is (numerically) an integer *)
|
||
rounded = Round[mat, 10^-8];
|
||
Print[title, " ", rounded ];
|
||
]
|
||
|
||
(* --------------------------------------------------------------- *)
|
||
(* Test data *)
|
||
(* --------------------------------------------------------------- *)
|
||
|
||
varA = {{1, 2}, {3, 4}};
|
||
varB = {{5, 6}, {7, 8}};
|
||
|
||
varC = {{1, 1, 1, 1},
|
||
{2, 4, 8, 16},
|
||
{3, 9, 27, 81},
|
||
{4, 16, 64, 256}};
|
||
|
||
varD = {{4, -3, 4/3, -1/4},
|
||
{-13/3, 19/4, -7/3, 11/24},
|
||
{3/2, -2, 7/6, -1/4},
|
||
{-1/6, 1/4, -1/6, 1/24}};
|
||
|
||
varE = {{1, 2, 3, 4},
|
||
{5, 6, 7, 8},
|
||
{9, 10, 11, 12},
|
||
{13, 14, 15, 16}};
|
||
|
||
varF = IdentityMatrix[4];
|
||
|
||
r = Sqrt[2]/2;
|
||
R = {{r, r}, {-r, r}};
|
||
|
||
(* --------------------------------------------------------------- *)
|
||
(* Run the examples – compare ordinary (`.`) and Strassen results *)
|
||
(* --------------------------------------------------------------- *)
|
||
|
||
intPrint["Regular multiply: ", Transpose[varA].Transpose[varB]];
|
||
intPrint["Strassen multiply: ", Strassen[Transpose[varA], Transpose[varB]]];
|
||
|
||
intPrint["Regular multiply: ", varC . varD];
|
||
intPrint["Strassen multiply: ", Strassen[varC, varD]];
|
||
|
||
intPrint["Regular multiply: ", varE . varF];
|
||
intPrint["Strassen multiply: ", Strassen[varE, varF]];
|
||
|
||
intPrint["Regular multiply: ", R . R];
|
||
intPrint["Strassen multiply: ", Strassen[R, R]];
|