Data update

This commit is contained in:
Ingy döt Net 2025-06-11 20:16:52 -04:00
parent 72eb4943cb
commit 4d5544505c
2347 changed files with 62432 additions and 16731 deletions

View file

@ -18,3 +18,5 @@ While practical implementations of Strassen's algorithm usually switch to standa
:* [[wp:Strassen algorithm|Wikipedia article]]
<br><br>

View file

@ -0,0 +1,248 @@
#include <iostream>
#include <vector>
#include <iomanip>
#include <cmath>
#include <sstream>
#include <stdexcept>
class Matrix {
public:
std::vector<std::vector<double>> data;
size_t rows;
size_t cols;
Matrix(const std::vector<std::vector<double>>& data) : data(data) {
rows = data.size();
cols = (rows > 0) ? data[0].size() : 0;
}
size_t getRows() const {
return rows;
}
size_t getCols() const {
return cols;
}
void validateDimensions(const Matrix& other) const {
if (getRows() != other.getRows() || getCols() != other.getCols()) {
throw std::runtime_error("Matrices must have the same dimensions.");
}
}
void validateMultiplication(const Matrix& other) const {
if (getCols() != other.getRows()) {
throw std::runtime_error("Cannot multiply these matrices.");
}
}
void validateSquarePowerOfTwo() const {
if (getRows() != getCols()) {
throw std::runtime_error("Matrix must be square.");
}
if (getRows() == 0 || (getRows() & (getRows() - 1)) != 0) {
throw std::runtime_error("Size of matrix must be a power of two.");
}
}
Matrix operator+(const Matrix& other) const {
validateDimensions(other);
std::vector<std::vector<double>> result_data(rows, std::vector<double>(cols));
for (size_t i = 0; i < rows; ++i) {
for (size_t j = 0; j < cols; ++j) {
result_data[i][j] = data[i][j] + other.data[i][j];
}
}
return Matrix(result_data);
}
Matrix operator-(const Matrix& other) const {
validateDimensions(other);
std::vector<std::vector<double>> result_data(rows, std::vector<double>(cols));
for (size_t i = 0; i < rows; ++i) {
for (size_t j = 0; j < cols; ++j) {
result_data[i][j] = data[i][j] - other.data[i][j];
}
}
return Matrix(result_data);
}
Matrix operator*(const Matrix& other) const {
validateMultiplication(other);
std::vector<std::vector<double>> result_data(rows, std::vector<double>(other.cols));
for (size_t i = 0; i < rows; ++i) {
for (size_t j = 0; j < other.cols; ++j) {
double sum = 0.0;
for (size_t k = 0; k < other.rows; ++k) {
sum += data[i][k] * other.data[k][j];
}
result_data[i][j] = sum;
}
}
return Matrix(result_data);
}
friend std::ostream& operator<<(std::ostream& os, const Matrix& matrix) {
for (const auto& row : matrix.data) {
os << "[";
for (size_t i = 0; i < row.size(); ++i) {
os << row[i];
if (i < row.size() - 1) {
os << ", ";
}
}
os << "]" << std::endl;
}
return os;
}
std::string toStringWithPrecision(size_t p) const {
std::stringstream ss;
ss << std::fixed << std::setprecision(p);
double pow = std::pow(10.0, p);
for (const auto& row : data) {
ss << "[";
for (size_t i = 0; i < row.size(); ++i) {
double r = std::round(row[i] * pow) / pow;
std::string formatted = ss.str();
ss.str(std::string());
ss << r;
formatted = ss.str();
if (formatted == "-0") {
ss.str(std::string());
ss << "0";
formatted = ss.str();
}
ss.str(std::string());
ss << formatted;
if (i < row.size() - 1) {
ss << ", ";
}
}
ss << "]" << std::endl;
}
return ss.str();
}
static std::array<std::array<size_t, 6>, 4> params(size_t r, size_t c) {
return {
{{{0, r, 0, c, 0, 0}},
{{0, r, c, 2 * c, 0, c}},
{{r, 2 * r, 0, c, r, 0}},
{{r, 2 * r, c, 2 * c, r, c}}}
};
}
std::array<Matrix, 4> toQuarters() const {
size_t r = getRows() / 2;
size_t c = getCols() / 2;
auto p = Matrix::params(r, c);
std::array<Matrix, 4> quarters = {
Matrix(std::vector<std::vector<double>>(r, std::vector<double>(c, 0.0))),
Matrix(std::vector<std::vector<double>>(r, std::vector<double>(c, 0.0))),
Matrix(std::vector<std::vector<double>>(r, std::vector<double>(c, 0.0))),
Matrix(std::vector<std::vector<double>>(r, std::vector<double>(c, 0.0)))
};
for (size_t k = 0; k < 4; ++k) {
std::vector<std::vector<double>> q_data(r, std::vector<double>(c));
for (size_t i = p[k][0]; i < p[k][1]; ++i) {
for (size_t j = p[k][2]; j < p[k][3]; ++j) {
q_data[i - p[k][4]][j - p[k][5]] = data[i][j];
}
}
quarters[k] = Matrix(q_data);
}
return quarters;
}
static Matrix fromQuarters(const std::array<Matrix, 4>& q) {
size_t r = q[0].getRows();
size_t c = q[0].getCols();
auto p = Matrix::params(r, c);
size_t rows = r * 2;
size_t cols = c * 2;
std::vector<std::vector<double>> m_data(rows, std::vector<double>(cols, 0.0));
for (size_t k = 0; k < 4; ++k) {
for (size_t i = p[k][0]; i < p[k][1]; ++i) {
for (size_t j = p[k][2]; j < p[k][3]; ++j) {
m_data[i][j] = q[k].data[i - p[k][4]][j - p[k][5]];
}
}
}
return Matrix(m_data);
}
Matrix strassen(const Matrix& other) const {
validateSquarePowerOfTwo();
other.validateSquarePowerOfTwo();
if (getRows() != other.getRows() || getCols() != other.getCols()) {
throw std::runtime_error("Matrices must be square and of equal size for Strassen multiplication.");
}
if (getRows() == 1) {
return *this * other;
}
auto qa = toQuarters();
auto qb = other.toQuarters();
Matrix p1 = (qa[1] - qa[3]).strassen(qb[2] + qb[3]);
Matrix p2 = (qa[0] + qa[3]).strassen(qb[0] + qb[3]);
Matrix p3 = (qa[0] - qa[2]).strassen(qb[0] + qb[1]);
Matrix p4 = (qa[0] + qa[1]).strassen(qb[3]);
Matrix p5 = qa[0].strassen(qb[1] - qb[3]);
Matrix p6 = qa[3].strassen(qb[2] - qb[0]);
Matrix p7 = (qa[2] + qa[3]).strassen(qb[0]);
std::array<Matrix, 4> q = {
Matrix(std::vector<std::vector<double>>(qa[0].getRows(), std::vector<double>(qa[0].getCols(), 0.0))),
Matrix(std::vector<std::vector<double>>(qa[0].getRows(), std::vector<double>(qa[0].getCols(), 0.0))),
Matrix(std::vector<std::vector<double>>(qa[0].getRows(), std::vector<double>(qa[0].getCols(), 0.0))),
Matrix(std::vector<std::vector<double>>(qa[0].getRows(), std::vector<double>(qa[0].getCols(), 0.0)))
};
q[0] = p1 + p2 - p4 + p6;
q[1] = p4 + p5;
q[2] = p6 + p7;
q[3] = p2 - p3 + p5 - p7;
return Matrix::fromQuarters(q);
}
};
int main() {
Matrix a({ {1.0, 2.0}, {3.0, 4.0} });
Matrix b({ {5.0, 6.0}, {7.0, 8.0} });
Matrix c({ {1.0, 1.0, 1.0, 1.0}, {2.0, 4.0, 8.0, 16.0}, {3.0, 9.0, 27.0, 81.0}, {4.0, 16.0, 64.0, 256.0} });
Matrix d({ {4.0, -3.0, 4.0 / 3.0, -1.0 / 4.0}, {-13.0 / 3.0, 19.0 / 4.0, -7.0 / 3.0, 11.0 / 24.0}, {3.0 / 2.0, -2.0, 7.0 / 6.0, -1.0 / 4.0}, {-1.0 / 6.0, 1.0 / 4.0, -1.0 / 6.0, 1.0 / 24.0} });
Matrix e({ {1.0, 2.0, 3.0, 4.0}, {5.0, 6.0, 7.0, 8.0}, {9.0, 10.0, 11.0, 12.0}, {13.0, 14.0, 15.0, 16.0} });
Matrix f({ {1.0, 0.0, 0.0, 0.0}, {0.0, 1.0, 0.0, 0.0}, {0.0, 0.0, 1.0, 0.0}, {0.0, 0.0, 0.0, 1.0} });
std::cout << "Using 'normal' matrix multiplication:" << std::endl;
std::cout << " a * b = " << a * b << std::endl;
std::cout << " c * d = " << (c * d).toStringWithPrecision(6) << std::endl;
std::cout << " e * f = " << e * f << std::endl;
std::cout << "\nUsing 'Strassen' matrix multiplication:" << std::endl;
std::cout << " a * b = " << a.strassen(b) << std::endl;
std::cout << " c * d = " << c.strassen(d).toStringWithPrecision(6) << std::endl;
std::cout << " e * f = " << e.strassen(f) << std::endl;
return 0;
}

View file

@ -0,0 +1,323 @@
using System;
using System.Collections.Generic;
using System.Linq;
using System.Text;
class Matrix
{
public List<List<double>> data;
public int rows;
public int cols;
public Matrix(List<List<double>> data)
{
this.data = data;
rows = data.Count;
cols = (rows > 0) ? data[0].Count : 0;
}
public int GetRows()
{
return rows;
}
public int GetCols()
{
return cols;
}
public void ValidateDimensions(Matrix other)
{
if (GetRows() != other.GetRows() || GetCols() != other.GetCols())
{
throw new InvalidOperationException("Matrices must have the same dimensions.");
}
}
public void ValidateMultiplication(Matrix other)
{
if (GetCols() != other.GetRows())
{
throw new InvalidOperationException("Cannot multiply these matrices.");
}
}
public void ValidateSquarePowerOfTwo()
{
if (GetRows() != GetCols())
{
throw new InvalidOperationException("Matrix must be square.");
}
if (GetRows() == 0 || (GetRows() & (GetRows() - 1)) != 0)
{
throw new InvalidOperationException("Size of matrix must be a power of two.");
}
}
public static Matrix operator +(Matrix a, Matrix b)
{
a.ValidateDimensions(b);
List<List<double>> resultData = new List<List<double>>();
for (int i = 0; i < a.rows; ++i)
{
List<double> row = new List<double>();
for (int j = 0; j < a.cols; ++j)
{
row.Add(a.data[i][j] + b.data[i][j]);
}
resultData.Add(row);
}
return new Matrix(resultData);
}
public static Matrix operator -(Matrix a, Matrix b)
{
a.ValidateDimensions(b);
List<List<double>> resultData = new List<List<double>>();
for (int i = 0; i < a.rows; ++i)
{
List<double> row = new List<double>();
for (int j = 0; j < a.cols; ++j)
{
row.Add(a.data[i][j] - b.data[i][j]);
}
resultData.Add(row);
}
return new Matrix(resultData);
}
public static Matrix operator *(Matrix a, Matrix b)
{
a.ValidateMultiplication(b);
List<List<double>> resultData = new List<List<double>>();
for (int i = 0; i < a.rows; ++i)
{
List<double> row = new List<double>();
for (int j = 0; j < b.cols; ++j)
{
double sum = 0.0;
for (int k = 0; k < b.rows; ++k)
{
sum += a.data[i][k] * b.data[k][j];
}
row.Add(sum);
}
resultData.Add(row);
}
return new Matrix(resultData);
}
public override string ToString()
{
StringBuilder sb = new StringBuilder();
foreach (var row in data)
{
sb.Append("[");
for (int i = 0; i < row.Count; ++i)
{
sb.Append(row[i]);
if (i < row.Count - 1)
{
sb.Append(", ");
}
}
sb.AppendLine("]");
}
return sb.ToString();
}
public string ToStringWithPrecision(int p)
{
StringBuilder sb = new StringBuilder();
double pow = Math.Pow(10.0, p);
foreach (var row in data)
{
sb.Append("[");
for (int i = 0; i < row.Count; ++i)
{
double r = Math.Round(row[i] * pow) / pow;
string formatted = r.ToString($"F{p}");
if (formatted == "-0" + (p > 0 ? "." + new string('0', p) : ""))
{
formatted = "0" + (p > 0 ? "." + new string('0', p) : "");
}
sb.Append(formatted);
if (i < row.Count - 1)
{
sb.Append(", ");
}
}
sb.AppendLine("]");
}
return sb.ToString();
}
private static int[,] GetParams(int r, int c)
{
return new int[,]
{
{0, r, 0, c, 0, 0},
{0, r, c, 2 * c, 0, c},
{r, 2 * r, 0, c, r, 0},
{r, 2 * r, c, 2 * c, r, c}
};
}
public Matrix[] ToQuarters()
{
int r = GetRows() / 2;
int c = GetCols() / 2;
int[,] p = GetParams(r, c);
Matrix[] quarters = new Matrix[4];
for (int k = 0; k < 4; ++k)
{
List<List<double>> qData = new List<List<double>>();
for (int i = 0; i < r; i++)
{
List<double> row = new List<double>();
for (int j = 0; j < c; j++)
{
row.Add(0.0);
}
qData.Add(row);
}
for (int i = p[k, 0]; i < p[k, 1]; ++i)
{
for (int j = p[k, 2]; j < p[k, 3]; ++j)
{
qData[i - p[k, 4]][j - p[k, 5]] = data[i][j];
}
}
quarters[k] = new Matrix(qData);
}
return quarters;
}
public static Matrix FromQuarters(Matrix[] q)
{
int r = q[0].GetRows();
int c = q[0].GetCols();
int[,] p = GetParams(r, c);
int rows = r * 2;
int cols = c * 2;
List<List<double>> mData = new List<List<double>>();
for (int i = 0; i < rows; i++)
{
List<double> row = new List<double>();
for (int j = 0; j < cols; j++)
{
row.Add(0.0);
}
mData.Add(row);
}
for (int k = 0; k < 4; ++k)
{
for (int i = p[k, 0]; i < p[k, 1]; ++i)
{
for (int j = p[k, 2]; j < p[k, 3]; ++j)
{
mData[i][j] = q[k].data[i - p[k, 4]][j - p[k, 5]];
}
}
}
return new Matrix(mData);
}
public Matrix Strassen(Matrix other)
{
ValidateSquarePowerOfTwo();
other.ValidateSquarePowerOfTwo();
if (GetRows() != other.GetRows() || GetCols() != other.GetCols())
{
throw new InvalidOperationException("Matrices must be square and of equal size for Strassen multiplication.");
}
if (GetRows() == 1)
{
return this * other;
}
Matrix[] qa = ToQuarters();
Matrix[] qb = other.ToQuarters();
Matrix p1 = (qa[1] - qa[3]).Strassen(qb[2] + qb[3]);
Matrix p2 = (qa[0] + qa[3]).Strassen(qb[0] + qb[3]);
Matrix p3 = (qa[0] - qa[2]).Strassen(qb[0] + qb[1]);
Matrix p4 = (qa[0] + qa[1]).Strassen(qb[3]);
Matrix p5 = qa[0].Strassen(qb[1] - qb[3]);
Matrix p6 = qa[3].Strassen(qb[2] - qb[0]);
Matrix p7 = (qa[2] + qa[3]).Strassen(qb[0]);
Matrix[] q = new Matrix[4];
q[0] = p1 + p2 - p4 + p6;
q[1] = p4 + p5;
q[2] = p6 + p7;
q[3] = p2 - p3 + p5 - p7;
return FromQuarters(q);
}
}
class Program
{
static void Main(string[] args)
{
Matrix a = new Matrix(new List<List<double>> { new List<double> { 1.0, 2.0 }, new List<double> { 3.0, 4.0 } });
Matrix b = new Matrix(new List<List<double>> { new List<double> { 5.0, 6.0 }, new List<double> { 7.0, 8.0 } });
Matrix c = new Matrix(new List<List<double>>
{
new List<double> { 1.0, 1.0, 1.0, 1.0 },
new List<double> { 2.0, 4.0, 8.0, 16.0 },
new List<double> { 3.0, 9.0, 27.0, 81.0 },
new List<double> { 4.0, 16.0, 64.0, 256.0 }
});
Matrix d = new Matrix(new List<List<double>>
{
new List<double> { 4.0, -3.0, 4.0 / 3.0, -1.0 / 4.0 },
new List<double> { -13.0 / 3.0, 19.0 / 4.0, -7.0 / 3.0, 11.0 / 24.0 },
new List<double> { 3.0 / 2.0, -2.0, 7.0 / 6.0, -1.0 / 4.0 },
new List<double> { -1.0 / 6.0, 1.0 / 4.0, -1.0 / 6.0, 1.0 / 24.0 }
});
Matrix e = new Matrix(new List<List<double>>
{
new List<double> { 1.0, 2.0, 3.0, 4.0 },
new List<double> { 5.0, 6.0, 7.0, 8.0 },
new List<double> { 9.0, 10.0, 11.0, 12.0 },
new List<double> { 13.0, 14.0, 15.0, 16.0 }
});
Matrix f = new Matrix(new List<List<double>>
{
new List<double> { 1.0, 0.0, 0.0, 0.0 },
new List<double> { 0.0, 1.0, 0.0, 0.0 },
new List<double> { 0.0, 0.0, 1.0, 0.0 },
new List<double> { 0.0, 0.0, 0.0, 1.0 }
});
Console.WriteLine("Using 'normal' matrix multiplication:");
Console.WriteLine($" a * b = {a * b}");
Console.WriteLine($" c * d = {(c * d).ToStringWithPrecision(6)}");
Console.WriteLine($" e * f = {e * f}");
Console.WriteLine("\nUsing 'Strassen' matrix multiplication:");
Console.WriteLine($" a * b = {a.Strassen(b)}");
Console.WriteLine($" c * d = {c.Strassen(d).ToStringWithPrecision(6)}");
Console.WriteLine($" e * f = {e.Strassen(f)}");
}
}

View file

@ -0,0 +1,281 @@
import java.util.ArrayList;
import java.util.Arrays;
import java.util.List;
class Matrix {
public List<List<Double>> data;
public int rows;
public int cols;
public Matrix(List<List<Double>> data) {
this.data = data;
rows = data.size();
cols = (rows > 0) ? data.get(0).size() : 0;
}
public int getRows() {
return rows;
}
public int getCols() {
return cols;
}
public void validateDimensions(Matrix other) {
if (getRows() != other.getRows() || getCols() != other.getCols()) {
throw new RuntimeException("Matrices must have the same dimensions.");
}
}
public void validateMultiplication(Matrix other) {
if (getCols() != other.getRows()) {
throw new RuntimeException("Cannot multiply these matrices.");
}
}
public void validateSquarePowerOfTwo() {
if (getRows() != getCols()) {
throw new RuntimeException("Matrix must be square.");
}
if (getRows() == 0 || (getRows() & (getRows() - 1)) != 0) {
throw new RuntimeException("Size of matrix must be a power of two.");
}
}
public Matrix add(Matrix other) {
validateDimensions(other);
List<List<Double>> resultData = new ArrayList<>();
for (int i = 0; i < rows; ++i) {
List<Double> row = new ArrayList<>();
for (int j = 0; j < cols; ++j) {
row.add(data.get(i).get(j) + other.data.get(i).get(j));
}
resultData.add(row);
}
return new Matrix(resultData);
}
public Matrix subtract(Matrix other) {
validateDimensions(other);
List<List<Double>> resultData = new ArrayList<>();
for (int i = 0; i < rows; ++i) {
List<Double> row = new ArrayList<>();
for (int j = 0; j < cols; ++j) {
row.add(data.get(i).get(j) - other.data.get(i).get(j));
}
resultData.add(row);
}
return new Matrix(resultData);
}
public Matrix multiply(Matrix other) {
validateMultiplication(other);
List<List<Double>> resultData = new ArrayList<>();
for (int i = 0; i < rows; ++i) {
List<Double> row = new ArrayList<>();
for (int j = 0; j < other.cols; ++j) {
double sum = 0.0;
for (int k = 0; k < other.rows; ++k) {
sum += data.get(i).get(k) * other.data.get(k).get(j);
}
row.add(sum);
}
resultData.add(row);
}
return new Matrix(resultData);
}
@Override
public String toString() {
StringBuilder sb = new StringBuilder();
for (List<Double> row : data) {
sb.append("[");
for (int i = 0; i < row.size(); ++i) {
sb.append(row.get(i));
if (i < row.size() - 1) {
sb.append(", ");
}
}
sb.append("]\n");
}
return sb.toString();
}
public String toStringWithPrecision(int p) {
StringBuilder sb = new StringBuilder();
double pow = Math.pow(10.0, p);
for (List<Double> row : data) {
sb.append("[");
for (int i = 0; i < row.size(); ++i) {
double r = Math.round(row.get(i) * pow) / pow;
String formatted = String.format("%." + p + "f", r);
if (formatted.equals("-0" + (p > 0 ? "." + "0".repeat(p) : ""))) {
formatted = "0" + (p > 0 ? "." + "0".repeat(p) : "");
}
sb.append(formatted);
if (i < row.size() - 1) {
sb.append(", ");
}
}
sb.append("]\n");
}
return sb.toString();
}
private static int[][] getParams(int r, int c) {
return new int[][] {
{0, r, 0, c, 0, 0},
{0, r, c, 2 * c, 0, c},
{r, 2 * r, 0, c, r, 0},
{r, 2 * r, c, 2 * c, r, c}
};
}
public Matrix[] toQuarters() {
int r = getRows() / 2;
int c = getCols() / 2;
int[][] p = getParams(r, c);
Matrix[] quarters = new Matrix[4];
for (int k = 0; k < 4; ++k) {
List<List<Double>> qData = new ArrayList<>();
for (int i = 0; i < r; i++) {
List<Double> row = new ArrayList<>();
for (int j = 0; j < c; j++) {
row.add(0.0);
}
qData.add(row);
}
for (int i = p[k][0]; i < p[k][1]; ++i) {
for (int j = p[k][2]; j < p[k][3]; ++j) {
qData.get(i - p[k][4]).set(j - p[k][5], data.get(i).get(j));
}
}
quarters[k] = new Matrix(qData);
}
return quarters;
}
public static Matrix fromQuarters(Matrix[] q) {
int r = q[0].getRows();
int c = q[0].getCols();
int[][] p = getParams(r, c);
int rows = r * 2;
int cols = c * 2;
List<List<Double>> mData = new ArrayList<>();
for (int i = 0; i < rows; i++) {
List<Double> row = new ArrayList<>();
for (int j = 0; j < cols; j++) {
row.add(0.0);
}
mData.add(row);
}
for (int k = 0; k < 4; ++k) {
for (int i = p[k][0]; i < p[k][1]; ++i) {
for (int j = p[k][2]; j < p[k][3]; ++j) {
mData.get(i).set(j, q[k].data.get(i - p[k][4]).get(j - p[k][5]));
}
}
}
return new Matrix(mData);
}
public Matrix strassen(Matrix other) {
validateSquarePowerOfTwo();
other.validateSquarePowerOfTwo();
if (getRows() != other.getRows() || getCols() != other.getCols()) {
throw new RuntimeException("Matrices must be square and of equal size for Strassen multiplication.");
}
if (getRows() == 1) {
return this.multiply(other);
}
Matrix[] qa = toQuarters();
Matrix[] qb = other.toQuarters();
Matrix p1 = qa[1].subtract(qa[3]).strassen(qb[2].add(qb[3]));
Matrix p2 = qa[0].add(qa[3]).strassen(qb[0].add(qb[3]));
Matrix p3 = qa[0].subtract(qa[2]).strassen(qb[0].add(qb[1]));
Matrix p4 = qa[0].add(qa[1]).strassen(qb[3]);
Matrix p5 = qa[0].strassen(qb[1].subtract(qb[3]));
Matrix p6 = qa[3].strassen(qb[2].subtract(qb[0]));
Matrix p7 = qa[2].add(qa[3]).strassen(qb[0]);
Matrix[] q = new Matrix[4];
q[0] = p1.add(p2).subtract(p4).add(p6);
q[1] = p4.add(p5);
q[2] = p6.add(p7);
q[3] = p2.subtract(p3).add(p5).subtract(p7);
return fromQuarters(q);
}
}
public class Main{
public static void main(String[] args) {
List<List<Double>> aData = new ArrayList<>();
aData.add(Arrays.asList(1.0, 2.0));
aData.add(Arrays.asList(3.0, 4.0));
Matrix a = new Matrix(aData);
List<List<Double>> bData = new ArrayList<>();
bData.add(Arrays.asList(5.0, 6.0));
bData.add(Arrays.asList(7.0, 8.0));
Matrix b = new Matrix(bData);
List<List<Double>> cData = new ArrayList<>();
cData.add(Arrays.asList(1.0, 1.0, 1.0, 1.0));
cData.add(Arrays.asList(2.0, 4.0, 8.0, 16.0));
cData.add(Arrays.asList(3.0, 9.0, 27.0, 81.0));
cData.add(Arrays.asList(4.0, 16.0, 64.0, 256.0));
Matrix c = new Matrix(cData);
List<List<Double>> dData = new ArrayList<>();
dData.add(Arrays.asList(4.0, -3.0, 4.0 / 3.0, -1.0 / 4.0));
dData.add(Arrays.asList(-13.0 / 3.0, 19.0 / 4.0, -7.0 / 3.0, 11.0 / 24.0));
dData.add(Arrays.asList(3.0 / 2.0, -2.0, 7.0 / 6.0, -1.0 / 4.0));
dData.add(Arrays.asList(-1.0 / 6.0, 1.0 / 4.0, -1.0 / 6.0, 1.0 / 24.0));
Matrix d = new Matrix(dData);
List<List<Double>> eData = new ArrayList<>();
eData.add(Arrays.asList(1.0, 2.0, 3.0, 4.0));
eData.add(Arrays.asList(5.0, 6.0, 7.0, 8.0));
eData.add(Arrays.asList(9.0, 10.0, 11.0, 12.0));
eData.add(Arrays.asList(13.0, 14.0, 15.0, 16.0));
Matrix e = new Matrix(eData);
List<List<Double>> fData = new ArrayList<>();
fData.add(Arrays.asList(1.0, 0.0, 0.0, 0.0));
fData.add(Arrays.asList(0.0, 1.0, 0.0, 0.0));
fData.add(Arrays.asList(0.0, 0.0, 1.0, 0.0));
fData.add(Arrays.asList(0.0, 0.0, 0.0, 1.0));
Matrix f = new Matrix(fData);
System.out.println("Using 'normal' matrix multiplication:");
System.out.println(" a * b = " + a.multiply(b));
System.out.println(" c * d = " + c.multiply(d).toStringWithPrecision(6));
System.out.println(" e * f = " + e.multiply(f));
System.out.println("\nUsing 'Strassen' matrix multiplication:");
System.out.println(" a * b = " + a.strassen(b));
System.out.println(" c * d = " + c.strassen(d).toStringWithPrecision(6));
System.out.println(" e * f = " + e.strassen(f));
}
}

View file

@ -0,0 +1,318 @@
/**
* Represents the dimensions of a matrix.
* @typedef {object} Shape
* @property {number} rows - Number of rows.
* @property {number} cols - Number of columns.
*/
/**
* A matrix implemented as a wrapper around a 2D array.
*/
class Matrix {
/**
* Creates a Matrix instance.
* @param {number[][]} data - A 2D array representing the matrix data.
*/
constructor(data = []) {
if (!Array.isArray(data) || (data.length > 0 && !Array.isArray(data[0]))) {
throw new Error("Matrix data must be a 2D array.");
}
// Basic check for consistent row lengths
if (data.length > 1) {
const firstLen = data[0].length;
if (!data.every(row => row.length === firstLen)) {
throw new Error("Matrix rows must have consistent lengths.");
}
}
this.data = data;
}
/**
* Gets the dimensions (shape) of the matrix.
* @returns {Shape} An object with rows and cols properties.
*/
get shape() {
const rows = this.data.length;
const cols = rows > 0 ? this.data[0].length : 0;
return { rows, cols };
}
/**
* Creates a new Matrix assembled from nested blocks of matrices.
* @param {Matrix[][]} blocks - A 2D array of Matrix objects.
* @returns {Matrix} A new Matrix assembled from the blocks.
* @static
*/
static block(blocks) {
const newMatrixData = [];
for (const hblock of blocks) {
if (!hblock || hblock.length === 0) continue;
const numRowsInBlock = hblock[0].shape.rows; // Assume consistent rows within a hblock
for (let i = 0; i < numRowsInBlock; i++) {
let newRow = [];
for (const matrix of hblock) {
if (matrix.data[i]) { // Check if row exists
newRow = newRow.concat(matrix.data[i]);
} else {
// Handle potential inconsistencies if needed, maybe throw error or fill?
console.warn("Inconsistent row count during block assembly");
}
}
newMatrixData.push(newRow);
}
}
return new Matrix(newMatrixData);
}
/**
* Performs naive matrix multiplication (dot product).
* @param {Matrix} b - The matrix to multiply with.
* @returns {Matrix} The resulting matrix product.
*/
dot(b) {
if (!(b instanceof Matrix)) {
throw new Error("Argument must be a Matrix instance.");
}
const aShape = this.shape;
const bShape = b.shape;
if (aShape.cols !== bShape.rows) {
throw new Error(`Matrices incompatible for multiplication: ${aShape.cols} cols != ${bShape.rows} rows`);
}
const resultData = [];
for (let i = 0; i < aShape.rows; i++) {
resultData[i] = [];
for (let j = 0; j < bShape.cols; j++) {
let sum = 0;
for (let k = 0; k < aShape.cols; k++) {
sum += this.data[i][k] * b.data[k][j];
}
resultData[i][j] = sum;
}
}
return new Matrix(resultData);
}
/**
* Multiplies this matrix by another matrix (using naive multiplication).
* Equivalent to Python's __matmul__.
* @param {Matrix} b - The matrix to multiply with.
* @returns {Matrix} The resulting matrix product.
*/
multiply(b) {
return this.dot(b);
}
/**
* Adds another matrix to this matrix.
* Equivalent to Python's __add__.
* @param {Matrix} b - The matrix to add.
* @returns {Matrix} The resulting matrix sum.
*/
add(b) {
if (!(b instanceof Matrix)) {
throw new Error("Argument must be a Matrix instance.");
}
const aShape = this.shape;
const bShape = b.shape;
if (aShape.rows !== bShape.rows || aShape.cols !== bShape.cols) {
throw new Error("Matrices must have the same shape for addition.");
}
const resultData = this.data.map((row, i) =>
row.map((val, j) => val + b.data[i][j])
);
return new Matrix(resultData);
}
/**
* Subtracts another matrix from this matrix.
* Equivalent to Python's __sub__.
* @param {Matrix} b - The matrix to subtract.
* @returns {Matrix} The resulting matrix difference.
*/
subtract(b) {
if (!(b instanceof Matrix)) {
throw new Error("Argument must be a Matrix instance.");
}
const aShape = this.shape;
const bShape = b.shape;
if (aShape.rows !== bShape.rows || aShape.cols !== bShape.cols) {
throw new Error("Matrices must have the same shape for subtraction.");
}
const resultData = this.data.map((row, i) =>
row.map((val, j) => val - b.data[i][j])
);
return new Matrix(resultData);
}
/**
* Helper function to slice the matrix data.
* @param {number} rowStart - Starting row index (inclusive).
* @param {number} rowEnd - Ending row index (exclusive).
* @param {number} colStart - Starting column index (inclusive).
* @param {number} colEnd - Ending column index (exclusive).
* @returns {Matrix} A new Matrix containing the sliced data.
* @private // Indicates intended internal use
*/
_slice(rowStart, rowEnd, colStart, colEnd) {
const slicedData = this.data.slice(rowStart, rowEnd)
.map(row => row.slice(colStart, colEnd));
return new Matrix(slicedData);
}
/**
* Performs matrix multiplication using Strassen's algorithm.
* Requires square matrices whose dimensions are powers of 2.
* @param {Matrix} b - The matrix to multiply with.
* @returns {Matrix} The resulting matrix product.
*/
strassen(b) {
if (!(b instanceof Matrix)) {
throw new Error("Argument must be a Matrix instance.");
}
const aShape = this.shape;
const bShape = b.shape;
if (aShape.rows !== aShape.cols) {
throw new Error("Matrix must be square for Strassen's algorithm.");
}
if (aShape.rows !== bShape.rows || aShape.cols !== bShape.cols) {
throw new Error("Matrices must have the same shape for Strassen's algorithm.");
}
// Check if dimension is a power of 2
if (aShape.rows === 0 || (aShape.rows & (aShape.rows - 1)) !== 0) {
throw new Error("Matrix dimension must be a power of 2 for Strassen's algorithm.");
}
if (aShape.rows === 1) {
return this.dot(b); // Base case
}
const n = aShape.rows;
const p = n / 2; // Partition size
// Partition matrices
const a11 = this._slice(0, p, 0, p);
const a12 = this._slice(0, p, p, n);
const a21 = this._slice(p, n, 0, p);
const a22 = this._slice(p, n, p, n);
const b11 = b._slice(0, p, 0, p);
const b12 = b._slice(0, p, p, n);
const b21 = b._slice(p, n, 0, p);
const b22 = b._slice(p, n, p, n);
// Recursive calls (Strassen's 7 multiplications)
const m1 = (a11.add(a22)).strassen(b11.add(b22));
const m2 = (a21.add(a22)).strassen(b11);
const m3 = a11.strassen(b12.subtract(b22));
const m4 = a22.strassen(b21.subtract(b11));
const m5 = (a11.add(a12)).strassen(b22);
const m6 = (a21.subtract(a11)).strassen(b11.add(b12));
const m7 = (a12.subtract(a22)).strassen(b21.add(b22));
// Combine results
const c11 = m1.add(m4).subtract(m5).add(m7);
const c12 = m3.add(m5);
const c21 = m2.add(m4);
const c22 = m1.subtract(m2).add(m3).add(m6);
// Assemble the final matrix from blocks
return Matrix.block([[c11, c12], [c21, c22]]);
}
/**
* Rounds the elements of the matrix to a specified number of decimal places.
* @param {number} [ndigits=0] - Number of decimal places to round to. If undefined or 0, rounds to the nearest integer.
* @returns {Matrix} A new Matrix with rounded elements.
*/
round(ndigits = 0) {
const factor = Math.pow(10, ndigits);
const roundFn = ndigits > 0
? (num) => Math.round((num + Number.EPSILON) * factor) / factor
: (num) => Math.round(num);
const roundedData = this.data.map(row =>
row.map(val => roundFn(val))
);
return new Matrix(roundedData);
}
/**
* Provides a string representation of the matrix.
* @returns {string} The string representation.
*/
toString() {
const rowsStr = this.data.map(row => ` [${row.join(', ')}]`);
return `Matrix([\n${rowsStr.join(',\n')}\n])`;
}
}
// --- Examples ---
function examples() {
const a = new Matrix([
[1, 2],
[3, 4],
]);
const b = new Matrix([
[5, 6],
[7, 8],
]);
const c = new Matrix([
[1, 1, 1, 1],
[2, 4, 8, 16],
[3, 9, 27, 81],
[4, 16, 64, 256],
]);
const d = new Matrix([
[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],
]);
const e = new Matrix([
[1, 2, 3, 4],
[5, 6, 7, 8],
[9, 10, 11, 12],
[13, 14, 15, 16],
]);
const f = new Matrix([ // Identity matrix
[1, 0, 0, 0],
[0, 1, 0, 0],
[0, 0, 1, 0],
[0, 0, 0, 1],
]);
console.log("Naive matrix multiplication:");
console.log(` a * b = ${a.multiply(b)}`); // Uses toString implicitly
console.log(` c * d = ${c.multiply(d).round(2)}`); // Round near-zero elements
console.log(` e * f = ${e.multiply(f)}`);
console.log("\nStrassen's matrix multiplication:");
console.log(` a * b = ${a.strassen(b)}`);
console.log(` c * d = ${c.strassen(d).round(2)}`); // Round near-zero elements
console.log(` e * f = ${e.strassen(f)}`);
// Example of addition/subtraction
console.log("\nAddition/Subtraction:");
const sum_ab = a.add(b);
console.log(` a + b = ${sum_ab}`);
const diff_ba = b.subtract(a);
console.log(` b - a = ${diff_ba}`);
// Example of block creation (creates a 4x4 matrix from four 2x2 matrices)
console.log("\nBlock Creation:");
const blocked = Matrix.block([[a, b], [b, a]]);
console.log(` Blocked [a,b],[b,a] = ${blocked}`);
}
// Run examples
examples();

View file

@ -0,0 +1,316 @@
class Matrix {
/** @type {number[][]} */
data;
/** @type {number} */
rows;
/** @type {number} */
cols;
/**
* @param {number[][]} data The matrix data as a 2D array.
*/
constructor(data) {
if (!Array.isArray(data) || (data.length > 0 && !Array.isArray(data[0]))) {
throw new Error("Input data must be a 2D array.");
}
// Optional: Deep copy to prevent external modifications
this.data = data.map(row => [...row]);
this.rows = data.length;
this.cols = (this.rows > 0) ? (data[0]?.length ?? 0) : 0; // Handle empty rows gracefully
// Optional: Validate that all rows have the same length
if (this.rows > 0) {
const firstRowLength = this.cols;
for (let i = 1; i < this.rows; i++) {
if (data[i].length !== firstRowLength) {
throw new Error("All rows in the matrix must have the same length.");
}
}
}
}
/** @returns {number} */
getRows() {
return this.rows;
}
/** @returns {number} */
getCols() {
return this.cols;
}
/** @param {Matrix} other */
validateDimensions(other) {
if (this.getRows() !== other.getRows() || this.getCols() !== other.getCols()) {
throw new Error("Matrices must have the same dimensions.");
}
}
/** @param {Matrix} other */
validateMultiplication(other) {
if (this.getCols() !== other.getRows()) {
throw new Error(`Cannot multiply matrices: (${this.getRows()}x${this.getCols()}) * (${other.getRows()}x${other.getCols()})`);
}
}
validateSquarePowerOfTwo() {
if (this.getRows() !== this.getCols()) {
throw new Error("Matrix must be square for this operation.");
}
const n = this.getRows();
// Check if n is 0 or not a power of two
// (n & (n - 1)) === 0 checks if n is a power of two (or 0)
if (n === 0 || (n & (n - 1)) !== 0) {
throw new Error("Size of matrix must be a power of two for Strassen.");
}
}
/**
* Adds another matrix to this matrix.
* @param {Matrix} other The matrix to add.
* @returns {Matrix} A new matrix representing the sum.
*/
add(other) {
this.validateDimensions(other);
const result_data = Array.from({ length: this.rows }, () => Array(this.cols).fill(0.0));
for (let i = 0; i < this.rows; ++i) {
for (let j = 0; j < this.cols; ++j) {
result_data[i][j] = this.data[i][j] + other.data[i][j];
}
}
return new Matrix(result_data);
}
/**
* Subtracts another matrix from this matrix.
* @param {Matrix} other The matrix to subtract.
* @returns {Matrix} A new matrix representing the difference.
*/
subtract(other) {
this.validateDimensions(other);
const result_data = Array.from({ length: this.rows }, () => Array(this.cols).fill(0.0));
for (let i = 0; i < this.rows; ++i) {
for (let j = 0; j < this.cols; ++j) {
result_data[i][j] = this.data[i][j] - other.data[i][j];
}
}
return new Matrix(result_data);
}
/**
* Multiplies this matrix by another matrix (standard algorithm).
* @param {Matrix} other The matrix to multiply by.
* @returns {Matrix} A new matrix representing the product.
*/
multiply(other) {
this.validateMultiplication(other);
const result_data = Array.from({ length: this.rows }, () => Array(other.cols).fill(0.0));
for (let i = 0; i < this.rows; ++i) {
for (let j = 0; j < other.cols; ++j) {
let sum = 0.0;
// K loops through columns of 'this' and rows of 'other'
for (let k = 0; k < this.cols; ++k) {
sum += this.data[i][k] * other.data[k][j];
}
result_data[i][j] = sum;
}
}
return new Matrix(result_data);
}
/**
* Returns a string representation of the matrix.
* @returns {string}
*/
toString() {
return this.data.map(row => `[${row.join(', ')}]`).join('\n');
}
/**
* Returns a string representation with specified precision, handling rounding and "-0".
* @param {number} p Precision (number of decimal places).
* @returns {string}
*/
toStringWithPrecision(p) {
let resultString = "";
const pow = Math.pow(10, p);
const zeroString = (0).toFixed(p);
const negZeroString = `-${zeroString}`;
for (const row of this.data) {
resultString += "[";
for (let i = 0; i < row.length; ++i) {
let val = row[i];
// Round like C++: round(val * 10^p) / 10^p
let roundedVal = Math.round(val * pow) / pow;
// Format to fixed precision
let formattedVal = roundedVal.toFixed(p);
// Handle the "-0.00..." case that toFixed might produce after rounding
if (formattedVal === negZeroString) {
formattedVal = zeroString;
}
resultString += formattedVal;
if (i < row.length - 1) {
resultString += ", ";
}
}
resultString += "]\n"; // Add newline after each row like C++ example
}
return resultString.trimEnd(); // Remove trailing newline
}
/**
* Helper function to get quadrant slicing parameters.
* @param {number} r Half rows
* @param {number} c Half columns
* @returns {number[][]} Array of [startRow, endRow, startCol, endCol, offsetRow, offsetCol]
*/
static params(r, c) {
// [startRow, endRow, startCol, endCol, resultOffsetRow, resultOffsetCol]
return [
[0, r, 0, c, 0, 0], // Top-left quadrant (0)
[0, r, c, 2 * c, 0, c], // Top-right quadrant (1)
[r, 2 * r, 0, c, r, 0], // Bottom-left quadrant (2)
[r, 2 * r, c, 2 * c, r, c] // Bottom-right quadrant (3)
];
}
/**
* Splits the matrix into four equally sized quadrants.
* Assumes matrix dimensions are even.
* @returns {Matrix[]} An array of four matrices [TopLeft, TopRight, BottomLeft, BottomRight].
*/
toQuarters() {
const r = this.getRows() / 2;
const c = this.getCols() / 2;
if (!Number.isInteger(r) || !Number.isInteger(c)) {
throw new Error("Matrix dimensions must be even for splitting into quarters.");
}
const p = Matrix.params(r, c);
const quarters = Array(4); // Will hold 4 Matrix objects
for (let k = 0; k < 4; ++k) {
const q_data = Array.from({ length: r }, () => Array(c));
const [startRow, endRow, startCol, endCol, offsetRow, offsetCol] = p[k];
for (let i = startRow; i < endRow; ++i) {
for (let j = startCol; j < endCol; ++j) {
// Adjust indices for the smaller quarter matrix
q_data[i - offsetRow][j - offsetCol] = this.data[i][j];
}
}
quarters[k] = new Matrix(q_data);
}
return quarters; // [TopLeft, TopRight, BottomLeft, BottomRight]
}
/**
* Creates a new matrix by combining four quadrant matrices.
* @param {Matrix[]} q An array of four matrices [TopLeft, TopRight, BottomLeft, BottomRight].
* @returns {Matrix} The combined matrix.
*/
static fromQuarters(q) {
if (q.length !== 4) throw new Error("Requires exactly four quadrant matrices.");
// Basic validation: Ensure quadrants have compatible dimensions
const r = q[0].getRows();
const c = q[0].getCols();
if (q[1].getRows() !== r || q[1].getCols() !== c ||
q[2].getRows() !== r || q[2].getCols() !== c ||
q[3].getRows() !== r || q[3].getCols() !== c) {
throw new Error("Quadrant matrices must have the same dimensions.");
}
const p = Matrix.params(r, c);
const rows = r * 2;
const cols = c * 2;
const m_data = Array.from({ length: rows }, () => Array(cols));
for (let k = 0; k < 4; ++k) {
const [startRow, endRow, startCol, endCol, offsetRow, offsetCol] = p[k];
for (let i = startRow; i < endRow; ++i) {
for (let j = startCol; j < endCol; ++j) {
// Adjust indices to read from the correct quadrant
m_data[i][j] = q[k].data[i - offsetRow][j - offsetCol];
}
}
}
return new Matrix(m_data);
}
/**
* Multiplies this matrix by another using Strassen's algorithm.
* Assumes both matrices are square and their size is a power of two.
* @param {Matrix} other The matrix to multiply by.
* @returns {Matrix} The resulting matrix product.
*/
strassen(other) {
this.validateSquarePowerOfTwo();
other.validateSquarePowerOfTwo();
if (this.getRows() !== other.getRows()) { // Columns already checked by validateSquarePowerOfTwo
throw new Error("Matrices must be square and of equal size for Strassen multiplication.");
}
// Base case: If the matrix is 1x1
if (this.getRows() === 1) {
// Use standard multiplication for the 1x1 case
return this.multiply(other);
}
// Split matrices into quarters
const qa = this.toQuarters(); // [a11, a12, a21, a22]
const qb = other.toQuarters(); // [b11, b12, b21, b22]
// Calculate the 7 products recursively (P1 to P7)
const p1 = (qa[1].subtract(qa[3])).strassen(qb[2].add(qb[3])); // p1 = (a12 - a22) * (b21 + b22)
const p2 = (qa[0].add(qa[3])).strassen(qb[0].add(qb[3])); // p2 = (a11 + a22) * (b11 + b22)
const p3 = (qa[0].subtract(qa[2])).strassen(qb[0].add(qb[1])); // p3 = (a11 - a21) * (b11 + b12)
const p4 = (qa[0].add(qa[1])).strassen(qb[3]); // p4 = (a11 + a12) * b22
const p5 = qa[0].strassen(qb[1].subtract(qb[3])); // p5 = a11 * (b12 - b22)
const p6 = qa[3].strassen(qb[2].subtract(qb[0])); // p6 = a22 * (b21 - b11)
const p7 = (qa[2].add(qa[3])).strassen(qb[0]); // p7 = (a21 + a22) * b11
// Calculate the result quarters (C11, C12, C21, C22)
const c11 = p1.add(p2).subtract(p4).add(p6);
const c12 = p4.add(p5);
const c21 = p6.add(p7);
const c22 = p2.subtract(p3).add(p5).subtract(p7);
// Combine the quarters into the result matrix
return Matrix.fromQuarters([c11, c12, c21, c22]);
}
}
// --- Main execution (equivalent to C++ main) ---
function main() {
const a = new Matrix([[1.0, 2.0], [3.0, 4.0]]);
const b = new Matrix([[5.0, 6.0], [7.0, 8.0]]);
const c = new Matrix([[1.0, 1.0, 1.0, 1.0], [2.0, 4.0, 8.0, 16.0], [3.0, 9.0, 27.0, 81.0], [4.0, 16.0, 64.0, 256.0]]);
const d = new Matrix([[4.0, -3.0, 4.0 / 3.0, -1.0 / 4.0], [-13.0 / 3.0, 19.0 / 4.0, -7.0 / 3.0, 11.0 / 24.0], [3.0 / 2.0, -2.0, 7.0 / 6.0, -1.0 / 4.0], [-1.0 / 6.0, 1.0 / 4.0, -1.0 / 6.0, 1.0 / 24.0]]);
const e = new Matrix([[1.0, 2.0, 3.0, 4.0], [5.0, 6.0, 7.0, 8.0], [9.0, 10.0, 11.0, 12.0], [13.0, 14.0, 15.0, 16.0]]);
const f = new Matrix([[1.0, 0.0, 0.0, 0.0], [0.0, 1.0, 0.0, 0.0], [0.0, 0.0, 1.0, 0.0], [0.0, 0.0, 0.0, 1.0]]); // Identity Matrix
console.log("Using 'normal' matrix multiplication:");
console.log(` a * b = \n${a.multiply(b).toString()}`);
console.log(`\n c * d = \n${c.multiply(d).toStringWithPrecision(6)}`); // Should be close to identity
console.log(`\n e * f = \n${e.multiply(f).toString()}`); // Should be e
console.log("\nUsing 'Strassen' matrix multiplication:");
try {
console.log(` a * b = \n${a.strassen(b).toString()}`);
console.log(`\n c * d = \n${c.strassen(d).toStringWithPrecision(6)}`); // Should be close to identity
console.log(`\n e * f = \n${e.strassen(f).toString()}`); // Should be e
} catch (error) {
console.error("Strassen multiplication failed:", error.message);
}
}
// Run the main function
main();

View file

@ -0,0 +1,269 @@
use std::fmt;
use std::ops::{Add, Mul, Sub};
#[derive(Debug, Clone)]
struct Matrix {
data: Vec<Vec<f64>>,
rows: usize,
cols: usize,
}
impl Matrix {
fn new(data: Vec<Vec<f64>>) -> Self {
let rows = data.len();
let cols = if rows > 0 { data[0].len() } else { 0 };
Matrix { data, rows, cols }
}
fn rows(&self) -> usize {
self.rows
}
fn cols(&self) -> usize {
self.cols
}
fn validate_dimensions(&self, other: &Matrix) {
if self.rows() != other.rows() || self.cols() != other.cols() {
panic!("Matrices must have the same dimensions.");
}
}
fn validate_multiplication(&self, other: &Matrix) {
if self.cols() != other.rows() {
panic!("Cannot multiply these matrices.");
}
}
fn validate_square_power_of_two(&self) {
if self.rows() != self.cols() {
panic!("Matrix must be square.");
}
if self.rows() == 0 || (self.rows() & (self.rows() - 1)) != 0 {
panic!("Size of matrix must be a power of two.");
}
}
}
impl Add for Matrix {
type Output = Self;
fn add(self, other: Self) -> Self {
self.validate_dimensions(&other);
let mut result_data = Vec::with_capacity(self.rows());
for i in 0..self.rows() {
let mut row = Vec::with_capacity(self.cols());
for j in 0..self.cols() {
row.push(self.data[i][j] + other.data[i][j]);
}
result_data.push(row);
}
Matrix::new(result_data)
}
}
impl Sub for Matrix {
type Output = Self;
fn sub(self, other: Self) -> Self {
self.validate_dimensions(&other);
let mut result_data = Vec::with_capacity(self.rows());
for i in 0..self.rows() {
let mut row = Vec::with_capacity(self.cols());
for j in 0..self.cols() {
row.push(self.data[i][j] - other.data[i][j]);
}
result_data.push(row);
}
Matrix::new(result_data)
}
}
impl Mul for Matrix {
type Output = Self;
fn mul(self, other: Self) -> Self {
self.validate_multiplication(&other);
let mut result_data = Vec::with_capacity(self.rows());
for i in 0..self.rows() {
let mut row = Vec::with_capacity(other.cols());
for j in 0..other.cols() {
let mut sum = 0.0;
for k in 0..other.rows() {
sum += self.data[i][k] * other.data[k][j];
}
row.push(sum);
}
result_data.push(row);
}
Matrix::new(result_data)
}
}
impl fmt::Display for Matrix {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let mut s = String::new();
for row in &self.data {
s.push_str(&format!("{:?}\n", row));
}
write!(f, "{}", s)
}
}
impl Matrix {
fn to_string_with_precision(&self, p: usize) -> String {
let mut s = String::new();
let pow = 10.0_f64.powi(p as i32);
for row in &self.data {
let mut t = Vec::new();
for &val in row {
let r = (val * pow).round() / pow;
let formatted = format!("{}", r);
if formatted == "-0" {
t.push("0".to_string());
} else {
t.push(formatted);
}
}
s.push_str(&format!("{:?}\n", t));
}
s
}
fn params(r: usize, c: usize) -> [[usize; 6]; 4] {
[
[0, r, 0, c, 0, 0],
[0, r, c, 2 * c, 0, c],
[r, 2 * r, 0, c, r, 0],
[r, 2 * r, c, 2 * c, r, c],
]
}
fn to_quarters(&self) -> [Matrix; 4] {
let r = self.rows() / 2;
let c = self.cols() / 2;
let p = Matrix::params(r, c);
let mut quarters: [Matrix; 4] = [
Matrix::new(vec![vec![0.0; c]; r]),
Matrix::new(vec![vec![0.0; c]; r]),
Matrix::new(vec![vec![0.0; c]; r]),
Matrix::new(vec![vec![0.0; c]; r]),
];
for k in 0..4 {
let mut q_data = Vec::with_capacity(r);
for i in p[k][0]..p[k][1] {
let mut row = Vec::with_capacity(c);
for j in p[k][2]..p[k][3] {
row.push(self.data[i][j]);
}
q_data.push(row);
}
quarters[k] = Matrix::new(q_data);
}
quarters
}
fn from_quarters(q: [Matrix; 4]) -> Matrix {
let r = q[0].rows();
let c = q[0].cols();
let p = Matrix::params(r, c);
let rows = r * 2;
let cols = c * 2;
let mut m_data = vec![vec![0.0; cols]; rows];
for k in 0..4 {
for i in p[k][0]..p[k][1] {
for j in p[k][2]..p[k][3] {
m_data[i][j] = q[k].data[i - p[k][4]][j - p[k][5]];
}
}
}
Matrix::new(m_data)
}
fn strassen(&self, other: Matrix) -> Matrix {
self.validate_square_power_of_two();
other.validate_square_power_of_two();
if self.rows() != other.rows() || self.cols() != other.cols() {
panic!("Matrices must be square and of equal size for Strassen multiplication.");
}
if self.rows() == 1 {
return self.clone() * other;
}
let qa = self.to_quarters();
let qb = other.to_quarters();
let p1 = (qa[1].clone() - qa[3].clone()).strassen(qb[2].clone() + qb[3].clone());
let p2 = (qa[0].clone() + qa[3].clone()).strassen(qb[0].clone() + qb[3].clone());
let p3 = (qa[0].clone() - qa[2].clone()).strassen(qb[0].clone() + qb[1].clone());
let p4 = (qa[0].clone() + qa[1].clone()).strassen(qb[3].clone());
let p5 = qa[0].clone().strassen(qb[1].clone() - qb[3].clone());
let p6 = qa[3].clone().strassen(qb[2].clone() - qb[0].clone());
let p7 = (qa[2].clone() + qa[3].clone()).strassen(qb[0].clone());
let mut q: [Matrix; 4] = [
Matrix::new(vec![vec![0.0; qa[0].cols()]; qa[0].rows()]),
Matrix::new(vec![vec![0.0; qa[0].cols()]; qa[0].rows()]),
Matrix::new(vec![vec![0.0; qa[0].cols()]; qa[0].rows()]),
Matrix::new(vec![vec![0.0; qa[0].cols()]; qa[0].rows()]),
];
q[0] = p1.clone() + p2.clone() - p4.clone() + p6.clone();
q[1] = p4 + p5.clone();
q[2] = p6 + p7.clone();
q[3] = p2 - p3.clone() + p5 - p7;
Matrix::from_quarters(q)
}
}
fn main() {
let a = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]);
let b = Matrix::new(vec![vec![5.0, 6.0], vec![7.0, 8.0]]);
let c = Matrix::new(vec![
vec![1.0, 1.0, 1.0, 1.0],
vec![2.0, 4.0, 8.0, 16.0],
vec![3.0, 9.0, 27.0, 81.0],
vec![4.0, 16.0, 64.0, 256.0],
]);
let d = Matrix::new(vec![
vec![4.0, -3.0, 4.0 / 3.0, -1.0 / 4.0],
vec![-13.0 / 3.0, 19.0 / 4.0, -7.0 / 3.0, 11.0 / 24.0],
vec![3.0 / 2.0, -2.0, 7.0 / 6.0, -1.0 / 4.0],
vec![-1.0 / 6.0, 1.0 / 4.0, -1.0 / 6.0, 1.0 / 24.0],
]);
let e = Matrix::new(vec![
vec![1.0, 2.0, 3.0, 4.0],
vec![5.0, 6.0, 7.0, 8.0],
vec![9.0, 10.0, 11.0, 12.0],
vec![13.0, 14.0, 15.0, 16.0],
]);
let f = Matrix::new(vec![
vec![1.0, 0.0, 0.0, 0.0],
vec![0.0, 1.0, 0.0, 0.0],
vec![0.0, 0.0, 1.0, 0.0],
vec![0.0, 0.0, 0.0, 1.0],
]);
println!("Using 'normal' matrix multiplication:");
println!(" a * b = {}", a.clone() * b.clone());
println!(" c * d = {}", (c.clone() * d.clone()).to_string_with_precision(6));
println!(" e * f = {}", e.clone() * f.clone());
println!("\nUsing 'Strassen' matrix multiplication:");
println!(" a * b = {}", a.strassen(b));
println!(" c * d = {}", c.strassen(d).to_string_with_precision(6));
println!(" e * f = {}", e.strassen(f));
}

View file

@ -0,0 +1,477 @@
const std = @import("std");
const fmt = std.fmt;
const ArrayList = std.ArrayList;
const Allocator = std.mem.Allocator;
const Matrix = struct {
data: ArrayList(ArrayList(f64)),
rows: usize,
cols: usize,
allocator: Allocator,
pub fn init(allocator: Allocator, data: ArrayList(ArrayList(f64))) !Matrix {
const rows = data.items.len;
const cols = if (rows > 0) data.items[0].items.len else 0;
return Matrix{
.data = data,
.rows = rows,
.cols = cols,
.allocator = allocator,
};
}
pub fn deinit(self: *Matrix) void {
for (self.data.items) |*row| {
row.deinit();
}
self.data.deinit();
}
pub fn clone(self: Matrix) !Matrix {
var new_data = ArrayList(ArrayList(f64)).init(self.allocator);
try new_data.ensureTotalCapacity(self.rows);
for (self.data.items) |row| {
var new_row = ArrayList(f64).init(self.allocator);
try new_row.ensureTotalCapacity(self.cols);
try new_row.appendSlice(row.items);
try new_data.append(new_row);
}
return Matrix{
.data = new_data,
.rows = self.rows,
.cols = self.cols,
.allocator = self.allocator,
};
}
pub fn validateDimensions(self: Matrix, other: Matrix) !void {
if (self.rows != other.rows or self.cols != other.cols) {
return error.DimensionMismatch;
}
}
pub fn validateMultiplication(self: Matrix, other: Matrix) !void {
if (self.cols != other.rows) {
return error.CannotMultiply;
}
}
pub fn validateSquarePowerOfTwo(self: Matrix) !void {
if (self.rows != self.cols) {
return error.NotSquare;
}
if (self.rows == 0 or (self.rows & (self.rows - 1)) != 0) {
return error.NotPowerOfTwo;
}
}
pub fn add(self: Matrix, other: Matrix) !Matrix {
try self.validateDimensions(other);
var result_data = ArrayList(ArrayList(f64)).init(self.allocator);
try result_data.ensureTotalCapacity(self.rows);
for (0..self.rows) |i| {
var row = ArrayList(f64).init(self.allocator);
try row.ensureTotalCapacity(self.cols);
for (0..self.cols) |j| {
try row.append(self.data.items[i].items[j] + other.data.items[i].items[j]);
}
try result_data.append(row);
}
return try Matrix.init(self.allocator, result_data);
}
pub fn sub(self: Matrix, other: Matrix) !Matrix {
try self.validateDimensions(other);
var result_data = ArrayList(ArrayList(f64)).init(self.allocator);
try result_data.ensureTotalCapacity(self.rows);
for (0..self.rows) |i| {
var row = ArrayList(f64).init(self.allocator);
try row.ensureTotalCapacity(self.cols);
for (0..self.cols) |j| {
try row.append(self.data.items[i].items[j] - other.data.items[i].items[j]);
}
try result_data.append(row);
}
return try Matrix.init(self.allocator, result_data);
}
pub fn mul(self: Matrix, other: Matrix) !Matrix {
try self.validateMultiplication(other);
var result_data = ArrayList(ArrayList(f64)).init(self.allocator);
try result_data.ensureTotalCapacity(self.rows);
for (0..self.rows) |i| {
var row = ArrayList(f64).init(self.allocator);
try row.ensureTotalCapacity(other.cols);
for (0..other.cols) |j| {
var sum: f64 = 0.0;
for (0..self.cols) |k| {
sum += self.data.items[i].items[k] * other.data.items[k].items[j];
}
try row.append(sum);
}
try result_data.append(row);
}
return try Matrix.init(self.allocator, result_data);
}
pub fn format(self: Matrix, comptime _: []const u8, _: fmt.FormatOptions, writer: anytype) !void {
for (self.data.items) |row| {
try writer.print("{any}\n", .{row.items});
}
}
pub fn toStringWithPrecision(self: Matrix, p: usize, allocator: Allocator) ![]u8 {
var output = ArrayList(u8).init(allocator);
defer output.deinit();
const pow = std.math.pow(f64, 10.0, @as(f64, @floatFromInt(p)));
for (self.data.items) |row| {
var formatted_row = ArrayList([]const u8).init(allocator);
defer {
for (formatted_row.items) |item| {
allocator.free(item);
}
formatted_row.deinit();
}
for (row.items) |val| {
const r = @round(val * pow) / pow;
const formatted = try fmt.allocPrint(allocator, "{d}", .{r});
if (std.mem.eql(u8, formatted, "-0")) {
allocator.free(formatted);
try formatted_row.append(try allocator.dupe(u8, "0"));
} else {
try formatted_row.append(formatted);
}
}
std.debug.print("{any}\n", .{formatted_row.items});
}
return output.toOwnedSlice();
}
fn params(r: usize, c: usize) [4][6]usize {
return [4][6]usize{
[_]usize{ 0, r, 0, c, 0, 0 },
[_]usize{ 0, r, c, 2 * c, 0, c },
[_]usize{ r, 2 * r, 0, c, r, 0 },
[_]usize{ r, 2 * r, c, 2 * c, r, c },
};
}
pub fn toQuarters(self: Matrix) ![4]Matrix {
const r = self.rows / 2;
const c = self.cols / 2;
const p = Matrix.params(r, c);
var quarters: [4]Matrix = undefined;
for (0..4) |k| {
var q_data = ArrayList(ArrayList(f64)).init(self.allocator);
try q_data.ensureTotalCapacity(r);
for (p[k][0]..p[k][1]) |i| {
var row = ArrayList(f64).init(self.allocator);
try row.ensureTotalCapacity(c);
for (p[k][2]..p[k][3]) |j| {
try row.append(self.data.items[i].items[j]);
}
try q_data.append(row);
}
quarters[k] = try Matrix.init(self.allocator, q_data);
}
return quarters;
}
pub fn fromQuarters(q: [4]Matrix, allocator: Allocator) !Matrix {
const r = q[0].rows;
const c = q[0].cols;
const p = Matrix.params(r, c);
const rows = r * 2;
const cols = c * 2;
var m_data = ArrayList(ArrayList(f64)).init(allocator);
try m_data.ensureTotalCapacity(rows);
for (0..rows) |_| {
var row = ArrayList(f64).init(allocator);
try row.ensureTotalCapacity(cols);
for (0..cols) |_| {
try row.append(0.0);
}
try m_data.append(row);
}
for (0..4) |k| {
for (p[k][0]..p[k][1]) |i| {
for (p[k][2]..p[k][3]) |j| {
m_data.items[i].items[j] = q[k].data.items[i - p[k][4]].items[j - p[k][5]];
}
}
}
return try Matrix.init(allocator, m_data);
}
pub fn strassen(self: Matrix, other: Matrix) !Matrix {
try self.validateSquarePowerOfTwo();
try other.validateSquarePowerOfTwo();
if (self.rows != other.rows or self.cols != other.cols) {
return error.InvalidDimensions;
}
if (self.rows == 1) {
return self.mul(other);
}
var qa = try self.toQuarters();
defer for (&qa) |*q| q.deinit();
var qb = try other.toQuarters();
defer for (&qb) |*q| q.deinit();
var t1 = try qa[1].sub(qa[3]);
defer t1.deinit();
var t2 = try qb[2].add(qb[3]);
defer t2.deinit();
var p1 = try t1.strassen(t2);
defer p1.deinit();
var t3 = try qa[0].add(qa[3]);
defer t3.deinit();
var t4 = try qb[0].add(qb[3]);
defer t4.deinit();
var p2 = try t3.strassen(t4);
defer p2.deinit();
var t5 = try qa[0].sub(qa[2]);
defer t5.deinit();
var t6 = try qb[0].add(qb[1]);
defer t6.deinit();
var p3 = try t5.strassen(t6);
defer p3.deinit();
var t7 = try qa[0].add(qa[1]);
defer t7.deinit();
var p4 = try t7.strassen(qb[3]);
defer p4.deinit();
var t8 = try qb[1].sub(qb[3]);
defer t8.deinit();
var p5 = try qa[0].strassen(t8);
defer p5.deinit();
var t9 = try qb[2].sub(qb[0]);
defer t9.deinit();
var p6 = try qa[3].strassen(t9);
defer p6.deinit();
var t10 = try qa[2].add(qa[3]);
defer t10.deinit();
var p7 = try t10.strassen(qb[0]);
defer p7.deinit();
var q: [4]Matrix = undefined;
// q[0] = p1 + p2 - p4 + p6
var ta = try p1.add(p2);
defer ta.deinit();
var tb = try ta.sub(p4);
defer tb.deinit();
q[0] = try tb.add(p6);
// q[1] = p4 + p5
q[1] = try p4.add(p5);
// q[2] = p6 + p7
q[2] = try p6.add(p7);
// q[3] = p2 - p3 + p5 - p7
var tc = try p2.sub(p3);
defer tc.deinit();
var td = try tc.add(p5);
defer td.deinit();
q[3] = try td.sub(p7);
defer for (&q) |*quarter| quarter.deinit();
return Matrix.fromQuarters(q, self.allocator);
}
};
pub fn main() !void {
var gpa = std.heap.GeneralPurposeAllocator(.{}){};
defer _ = gpa.deinit();
const allocator = gpa.allocator();
// Matrix A - [1 2; 3 4]
var a_data = ArrayList(ArrayList(f64)).init(allocator);
var a_row1 = ArrayList(f64).init(allocator);
try a_row1.appendSlice(&[_]f64{ 1.0, 2.0 });
var a_row2 = ArrayList(f64).init(allocator);
try a_row2.appendSlice(&[_]f64{ 3.0, 4.0 });
try a_data.append(a_row1);
try a_data.append(a_row2);
var a = try Matrix.init(allocator, a_data);
defer a.deinit();
// Matrix B - [5 6; 7 8]
var b_data = ArrayList(ArrayList(f64)).init(allocator);
var b_row1 = ArrayList(f64).init(allocator);
try b_row1.appendSlice(&[_]f64{ 5.0, 6.0 });
var b_row2 = ArrayList(f64).init(allocator);
try b_row2.appendSlice(&[_]f64{ 7.0, 8.0 });
try b_data.append(b_row1);
try b_data.append(b_row2);
var b = try Matrix.init(allocator, b_data);
defer b.deinit();
// Matrix C - 4x4
var c_data = ArrayList(ArrayList(f64)).init(allocator);
var c_row1 = ArrayList(f64).init(allocator);
try c_row1.appendSlice(&[_]f64{ 1.0, 1.0, 1.0, 1.0 });
var c_row2 = ArrayList(f64).init(allocator);
try c_row2.appendSlice(&[_]f64{ 2.0, 4.0, 8.0, 16.0 });
var c_row3 = ArrayList(f64).init(allocator);
try c_row3.appendSlice(&[_]f64{ 3.0, 9.0, 27.0, 81.0 });
var c_row4 = ArrayList(f64).init(allocator);
try c_row4.appendSlice(&[_]f64{ 4.0, 16.0, 64.0, 256.0 });
try c_data.append(c_row1);
try c_data.append(c_row2);
try c_data.append(c_row3);
try c_data.append(c_row4);
var c = try Matrix.init(allocator, c_data);
defer c.deinit();
// Matrix D - 4x4
var d_data = ArrayList(ArrayList(f64)).init(allocator);
var d_row1 = ArrayList(f64).init(allocator);
try d_row1.appendSlice(&[_]f64{ 4.0, -3.0, 4.0 / 3.0, -1.0 / 4.0 });
var d_row2 = ArrayList(f64).init(allocator);
try d_row2.appendSlice(&[_]f64{ -13.0 / 3.0, 19.0 / 4.0, -7.0 / 3.0, 11.0 / 24.0 });
var d_row3 = ArrayList(f64).init(allocator);
try d_row3.appendSlice(&[_]f64{ 3.0 / 2.0, -2.0, 7.0 / 6.0, -1.0 / 4.0 });
var d_row4 = ArrayList(f64).init(allocator);
try d_row4.appendSlice(&[_]f64{ -1.0 / 6.0, 1.0 / 4.0, -1.0 / 6.0, 1.0 / 24.0 });
try d_data.append(d_row1);
try d_data.append(d_row2);
try d_data.append(d_row3);
try d_data.append(d_row4);
var d = try Matrix.init(allocator, d_data);
defer d.deinit();
// Matrix E - 4x4
var e_data = ArrayList(ArrayList(f64)).init(allocator);
var e_row1 = ArrayList(f64).init(allocator);
try e_row1.appendSlice(&[_]f64{ 1.0, 2.0, 3.0, 4.0 });
var e_row2 = ArrayList(f64).init(allocator);
try e_row2.appendSlice(&[_]f64{ 5.0, 6.0, 7.0, 8.0 });
var e_row3 = ArrayList(f64).init(allocator);
try e_row3.appendSlice(&[_]f64{ 9.0, 10.0, 11.0, 12.0 });
var e_row4 = ArrayList(f64).init(allocator);
try e_row4.appendSlice(&[_]f64{ 13.0, 14.0, 15.0, 16.0 });
try e_data.append(e_row1);
try e_data.append(e_row2);
try e_data.append(e_row3);
try e_data.append(e_row4);
var e = try Matrix.init(allocator, e_data);
defer e.deinit();
// Matrix F - Identity 4x4
var f_data = ArrayList(ArrayList(f64)).init(allocator);
var f_row1 = ArrayList(f64).init(allocator);
try f_row1.appendSlice(&[_]f64{ 1.0, 0.0, 0.0, 0.0 });
var f_row2 = ArrayList(f64).init(allocator);
try f_row2.appendSlice(&[_]f64{ 0.0, 1.0, 0.0, 0.0 });
var f_row3 = ArrayList(f64).init(allocator);
try f_row3.appendSlice(&[_]f64{ 0.0, 0.0, 1.0, 0.0 });
var f_row4 = ArrayList(f64).init(allocator);
try f_row4.appendSlice(&[_]f64{ 0.0, 0.0, 0.0, 1.0 });
try f_data.append(f_row1);
try f_data.append(f_row2);
try f_data.append(f_row3);
try f_data.append(f_row4);
var f = try Matrix.init(allocator, f_data);
defer f.deinit();
const stdout = std.io.getStdOut().writer();
try stdout.print("Using 'normal' matrix multiplication:\n", .{});
var a_clone = try a.clone();
defer a_clone.deinit();
var b_clone = try b.clone();
defer b_clone.deinit();
var ab = try a_clone.mul(b_clone);
defer ab.deinit();
try stdout.print(" a * b = {}\n", .{ab});
var c_clone = try c.clone();
defer c_clone.deinit();
var d_clone = try d.clone();
defer d_clone.deinit();
var cd = try c_clone.mul(d_clone);
defer cd.deinit();
const cd_str = try cd.toStringWithPrecision(6, allocator);
defer allocator.free(cd_str);
try stdout.print(" c * d = {s}\n", .{cd_str});
var e_clone = try e.clone();
defer e_clone.deinit();
var f_clone = try f.clone();
defer f_clone.deinit();
var ef = try e_clone.mul(f_clone);
defer ef.deinit();
try stdout.print(" e * f = {}\n", .{ef});
try stdout.print("\nUsing 'Strassen' matrix multiplication:\n", .{});
var a_clone2 = try a.clone();
defer a_clone2.deinit();
var b_clone2 = try b.clone();
defer b_clone2.deinit();
var ab_s = try a_clone2.strassen(b_clone2);
defer ab_s.deinit();
try stdout.print(" a * b = {}\n", .{ab_s});
var c_clone2 = try c.clone();
defer c_clone2.deinit();
var d_clone2 = try d.clone();
defer d_clone2.deinit();
var cd_s = try c_clone2.strassen(d_clone2);
defer cd_s.deinit();
const cd_s_str = try cd_s.toStringWithPrecision(6, allocator);
defer allocator.free(cd_s_str);
try stdout.print(" c * d = {s}\n", .{cd_s_str});
var e_clone2 = try e.clone();
defer e_clone2.deinit();
var f_clone2 = try f.clone();
defer f_clone2.deinit();
var ef_s = try e_clone2.strassen(f_clone2);
defer ef_s.deinit();
try stdout.print(" e * f = {}\n", .{ef_s});
}