Data update
This commit is contained in:
parent
4bb20c9b71
commit
cbaf4c4b64
12390 changed files with 318560 additions and 27248 deletions
70
Task/Strassens-algorithm/Pluto/strassens-algorithm.pluto
Normal file
70
Task/Strassens-algorithm/Pluto/strassens-algorithm.pluto
Normal file
|
|
@ -0,0 +1,70 @@
|
|||
require "matrix"
|
||||
require "table2"
|
||||
local fmt = require "fmt"
|
||||
|
||||
local function strassen(a, b)
|
||||
assert(a:issquare() and a:samesize(b), "matrices must be square and of equal size")
|
||||
local n = a:numrows()
|
||||
assert(n & (n - 1) == 0, "size of matrices must be a power of two")
|
||||
if n == 1 then return a:matmul(b) end
|
||||
local h = n // 2
|
||||
|
||||
local a11 = a:submatrix(1, h, 1, h)
|
||||
local a12 = a:submatrix(1, h, h + 1, n)
|
||||
local a21 = a:submatrix(h + 1, n, 1, h)
|
||||
local a22 = a:submatrix(h + 1, n, h + 1, n)
|
||||
|
||||
local b11 = b:submatrix(1, h, 1, h)
|
||||
local b12 = b:submatrix(1, h, h + 1, n)
|
||||
local b21 = b:submatrix(h + 1, n, 1, h)
|
||||
local b22 = b:submatrix(h + 1, n, h + 1, n)
|
||||
|
||||
local m1 = strassen(a11 + a22, b11 + b22)
|
||||
local m2 = strassen(a21 + a22, b11 )
|
||||
local m3 = strassen(a11, b12 - b22)
|
||||
local m4 = strassen(a22, b21 - b11)
|
||||
local m5 = strassen(a11 + a12, b22 )
|
||||
local m6 = strassen(a21 - a11, b11 + b12)
|
||||
local m7 = strassen(a12 - a22, b21 + b22)
|
||||
|
||||
local c11 = m1 + m4 - m5 + m7
|
||||
local c12 = m3 + m5
|
||||
local c21 = m2 + m4
|
||||
local c22 = m1 - m2 + m3 + m6
|
||||
|
||||
local c = {}
|
||||
for i = 1, n do c[i] = table.rep(n, 0) end
|
||||
for i = 1, h do
|
||||
for j = 1, h do
|
||||
c[i][j] = c11:get(i, j)
|
||||
c[i][j + h] = c12:get(i, j)
|
||||
c[i + h][j] = c21:get(i, j)
|
||||
c[i + h][j + h] = c22:get(i, j)
|
||||
end
|
||||
end
|
||||
return matrix.from(c)
|
||||
end
|
||||
|
||||
-- Round numbers which differ from an integer by less than 'tol' to that integer.
|
||||
local function round(n, tol = 1e-12)
|
||||
if math.abs(n - math.round(n)) < tol then return math.round(n) end
|
||||
return n
|
||||
end
|
||||
|
||||
local a = matrix.from({ {1, 2}, {3, 4} })
|
||||
local b = matrix.from({ {5, 6}, {7, 8} })
|
||||
local c = matrix.from({ {1, 1, 1, 1}, {2, 4, 8, 16}, {3, 9, 27, 81}, {4, 16, 64, 256} })
|
||||
local d = matrix.from({ {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} })
|
||||
local e = matrix.from({ {1, 2, 3, 4}, {5, 6, 7, 8}, {9, 10, 11, 12}, {13, 14, 15, 16} })
|
||||
local f = matrix.from({ {1, 0, 0, 0}, {0, 1, 0, 0}, {0, 0, 1, 0}, {0, 0, 0, 1} })
|
||||
|
||||
print("Using 'normal' matrix multiplication:")
|
||||
fmt.print(" a x b = %s", fmt.swrite(a:matmul(b):toarray()))
|
||||
fmt.print(" c x d = %s", fmt.swrite(c:matmul(d):map(round):toarray()))
|
||||
fmt.print(" e x f = %s", fmt.swrite(e:matmul(f):toarray()))
|
||||
|
||||
print("\nUsing 'Strassen' matrix multiplication:")
|
||||
fmt.print(" a x b = %s", fmt.swrite(strassen(a, b):toarray()))
|
||||
fmt.print(" c x d = %s", fmt.swrite(strassen(c, d):map(round):toarray()))
|
||||
fmt.print(" e x f = %s", fmt.swrite(strassen(e, f):toarray()))
|
||||
Loading…
Add table
Add a link
Reference in a new issue