最后一个样例不过qwq,求助daolao
查看原帖
最后一个样例不过qwq,求助daolao
1005061
Datieguy楼主2023/8/24 15:08
#include <iostream>
#include <vector>

using namespace std;

// 矩阵相加
vector<vector<int>> matrixAdd(const vector<vector<int>>& A, const vector<vector<int>>& B) 
{
    int n = A.size();
    int m = A[0].size();
    vector<vector<int>> C(n, vector<int>(m));

    for (int i = 0; i < n; i++) 
    {
        for (int j = 0; j < m; j++) 
        {
            C[i][j] = A[i][j] + B[i][j];
        }
    }

    return C;
}

// 矩阵相减
vector<vector<int>> matrixSubtract(const vector<vector<int>>& A, const vector<vector<int>>& B) {
    int n = A.size();
    int m = A[0].size();
    vector<vector<int>> C(n, vector<int>(m));

    for (int i = 0; i < n; i++) 
    {
        for (int j = 0; j < m; j++) 
        {
            C[i][j] = A[i][j] - B[i][j];
        }
    }

    return C;
}

// 普通矩阵乘法
vector<vector<int>> matrixMultiply(const vector<vector<int>>& A, const vector<vector<int>>& B) 
{
    int n = A.size();
    int m = A[0].size();
    int k = B[0].size();
    vector<vector<int>> C(n, vector<int>(k, 0));

    for (int i = 0; i < n; i++) 
    {
        for (int j = 0; j < k; j++) 
        {
            for (int l = 0; l < m; l++) 
            {
                C[i][j] += A[i][l] * B[l][j];
            }
        }
    }

    return C;
}

// Strassen矩阵乘法
vector<vector<int>> strassenMultiply(const vector<vector<int>>& A, const vector<vector<int>>& B) 
{
    int n = A.size();
    int m = A[0].size();
    int k = B[0].size();

    // 判断是否需要使用普通矩阵乘法
    if ( n % 2 != 0 || m != n || k != n) 
    {
        return matrixMultiply(A, B);
    }

    int halfN = n / 2;
    int halfK = k / 2;

    // 将原始矩阵划分成四个子矩阵
    vector<vector<int>> A11(halfN, vector<int>(halfK));
    vector<vector<int>> A12(halfN, vector<int>(halfK));
    vector<vector<int>> A21(halfN, vector<int>(halfK));
    vector<vector<int>> A22(halfN, vector<int>(halfK));
    vector<vector<int>> B11(halfN, vector<int>(halfK));
    vector<vector<int>> B12(halfN, vector<int>(halfK));
    vector<vector<int>> B21(halfN, vector<int>(halfK));
    vector<vector<int>> B22(halfN, vector<int>(halfK));

    // 将原始矩阵的值赋给子矩阵
    for (int i = 0; i < halfN; i++) 
    {
        for (int j = 0; j < halfK; j++) 
        {
            A11[i][j] = A[i][j];
            A12[i][j] = A[i][j + halfK];
            A21[i][j] = A[i + halfN][j];
            A22[i][j] = A[i + halfN][j + halfK];
            B11[i][j] = B[i][j];
            B12[i][j] = B[i][j + halfK];
            B21[i][j] = B[i + halfN][j];
            B22[i][j] = B[i + halfN][j + halfK];
        }
    }

    // 使用Strassen算法计算结果矩阵的四个子矩阵
    vector<vector<int>> P1 = strassenMultiply(matrixAdd(A11, A22), matrixAdd(B11, B22));
    vector<vector<int>> P2 = strassenMultiply(matrixAdd(A21, A22), B11);
    vector<vector<int>> P3 = strassenMultiply(A11, matrixSubtract(B12,B22));
    vector<vector<int>> P4 = strassenMultiply(A22, matrixSubtract(B21, B11));
    vector<vector<int>> P5 = strassenMultiply(matrixAdd(A11, A12), B22);
    vector<vector<int>> P6 = strassenMultiply(matrixSubtract(A21, A11), matrixAdd(B11, B12));
    vector<vector<int>> P7 = strassenMultiply(matrixSubtract(A12, A22), matrixAdd(B21, B22));

    // 计算结果矩阵的四个子矩阵
    vector<vector<int>> C11 = matrixAdd(matrixSubtract(matrixAdd(P1, P4), P5), P7);
    vector<vector<int>> C12 = matrixAdd(P3, P5);
    vector<vector<int>> C21 = matrixAdd(P2, P4);
    vector<vector<int>> C22 = matrixSubtract(matrixSubtract(matrixAdd(P1, P3), P2), P6);

    // 构建结果矩阵
    vector<vector<int>> C(n, vector<int>(k));

    // 将子矩阵的值赋给结果矩阵
    for (int i = 0; i < halfN; i++) 
    {
        for (int j = 0; j < halfK; j++) 
        {
            C[i][j] = C11[i][j];
            C[i][j + halfK] = C12[i][j];
            C[i + halfN][j] = C21[i][j];
            C[i + halfN][j + halfK] = C22[i][j];
        }
    }

    return C;
}

int main() 
{
    // 输入矩阵 A 和 B
    int n, m, k;
    cin >> n >> m >> k;

    vector<vector<int>> A(n, vector<int>(m));
    vector<vector<int>> B(m, vector<int>(k));

    for (int i = 0; i < n; i++) 
    {
        for (int j = 0; j < m; j++) 
        {
            cin >> A[i][j];
        }
    }

    for (int i = 0; i < m; i++) 
    {
        for (int j = 0; j < k; j++) 
        {
            cin >> B[i][j];
        }
    }

    // 调用 Strassen 矩阵乘法
    vector<vector<int>> C = strassenMultiply(A, B);

    // 输出结果矩阵 C
    for (int i = 0; i < n; i++) 
    {
        for (int j = 0; j < k; j++) 
        {
            cout << C[i][j] << " ";
        }
        cout << endl;
    }

    return 0;
}
2023/8/24 15:08
加载中...