Actual source code: mpimattransposematmult.c
1: /*
2: Defines matrix-matrix product routines for pairs of MPIAIJ matrices
3: C = A^T * B
4: The routines are slightly modified from MatTransposeMatMultxxx_SeqAIJ_SeqDense().
5: */
6: #include <../src/mat/impls/aij/seq/aij.h>
7: #include <../src/mat/impls/aij/mpi/mpiaij.h>
8: #include <../src/mat/impls/dense/mpi/mpidense.h>
10: static PetscErrorCode MatProductCtxDestroy_MPIDense_MatTransMatMult(PetscCtxRt data)
11: {
12: MatProductCtx_MatTransMatMult *atb = *(MatProductCtx_MatTransMatMult **)data;
14: PetscFunctionBegin;
15: PetscCall(MatDestroy(&atb->mA));
16: PetscCall(VecDestroy(&atb->bt));
17: PetscCall(VecDestroy(&atb->ct));
18: PetscCall(PetscFree(atb));
19: PetscFunctionReturn(PETSC_SUCCESS);
20: }
22: static PetscErrorCode MatTransposeMatMultNumeric_MPIAIJ_MPIDense(Mat, Mat, Mat);
24: PETSC_INTERN PetscErrorCode MatTransposeMatMultSymbolic_MPIAIJ_MPIDense(Mat A, Mat B, PetscReal fill, Mat C)
25: {
26: MatProductCtx_MatTransMatMult *atb;
27: PetscBool cisdense;
29: PetscFunctionBegin;
30: MatCheckProduct(C, 4);
31: PetscCheck(!C->product->data, PETSC_COMM_SELF, PETSC_ERR_PLIB, "Extra product struct not empty");
33: /* create output dense matrix C = A^T*B */
34: PetscCall(MatSetSizes(C, A->cmap->n, B->cmap->n, A->cmap->N, B->cmap->N));
35: PetscCall(PetscObjectTypeCompareAny((PetscObject)C, &cisdense, MATMPIDENSE, MATMPIDENSECUDA, ""));
36: if (!cisdense) {
37: PetscCall(MatSetType(C, ((PetscObject)B)->type_name));
38: PetscCall(MatSetVecType(C, B->defaultvectype));
39: }
40: PetscCall(MatSetUp(C));
42: /* create additional data structure for the product */
43: PetscCall(PetscNew(&atb));
44: if (B->cmap->N) {
45: PetscCall(MatCreateMAIJ(A, B->cmap->N, &atb->mA));
46: if (!atb->mA->assembled) {
47: PetscCall(MatAssemblyBegin(atb->mA, MAT_FINAL_ASSEMBLY));
48: PetscCall(MatAssemblyEnd(atb->mA, MAT_FINAL_ASSEMBLY));
49: }
50: PetscCall(MatCreateVecs(atb->mA, &atb->ct, &atb->bt));
51: }
52: C->product->data = atb;
53: C->product->destroy = MatProductCtxDestroy_MPIDense_MatTransMatMult;
55: C->ops->transposematmultnumeric = MatTransposeMatMultNumeric_MPIAIJ_MPIDense;
56: PetscFunctionReturn(PETSC_SUCCESS);
57: }
59: static PetscErrorCode MatTransposeMatMultNumeric_MPIAIJ_MPIDense(Mat A, Mat B, Mat C)
60: {
61: const PetscScalar *Barray, *ctarray;
62: PetscScalar *Carray, *btarray;
63: PetscInt i, j, m = A->rmap->n, n = A->cmap->n, ldb, BN = B->cmap->N, ldc;
64: MatProductCtx_MatTransMatMult *atb;
65: Vec bt, ct;
67: PetscFunctionBegin;
68: MatCheckProduct(C, 3);
69: atb = (MatProductCtx_MatTransMatMult *)C->product->data;
70: PetscCheck(atb, PETSC_COMM_SELF, PETSC_ERR_PLIB, "Missing product struct");
71: if (!BN) {
72: PetscCall(MatAssemblyBegin(C, MAT_FINAL_ASSEMBLY));
73: PetscCall(MatAssemblyEnd(C, MAT_FINAL_ASSEMBLY));
74: PetscFunctionReturn(PETSC_SUCCESS);
75: }
76: bt = atb->bt;
77: ct = atb->ct;
79: /* transpose local array of B, then copy it to vector bt */
80: PetscCall(MatDenseGetArrayRead(B, &Barray));
81: PetscCall(MatDenseGetLDA(B, &ldb));
82: PetscCall(VecGetArray(bt, &btarray));
83: for (j = 0; j < BN; j++)
84: for (i = 0; i < m; i++) btarray[i * BN + j] = Barray[j * ldb + i];
85: PetscCall(VecRestoreArray(bt, &btarray));
86: PetscCall(MatDenseRestoreArrayRead(B, &Barray));
88: /* compute ct = mA^T * cb */
89: PetscCall(MatMultTranspose(atb->mA, bt, ct));
91: /* transpose local array of ct to matrix C */
92: PetscCall(MatDenseGetArray(C, &Carray));
93: PetscCall(MatDenseGetLDA(C, &ldc));
94: PetscCall(VecGetArrayRead(ct, &ctarray));
95: for (j = 0; j < BN; j++)
96: for (i = 0; i < n; i++) Carray[j * ldc + i] = ctarray[i * BN + j];
97: PetscCall(VecRestoreArrayRead(ct, &ctarray));
98: PetscCall(MatDenseRestoreArray(C, &Carray));
99: PetscCall(MatSetOption(C, MAT_NO_OFF_PROC_ENTRIES, PETSC_TRUE));
100: PetscCall(MatAssemblyBegin(C, MAT_FINAL_ASSEMBLY));
101: PetscCall(MatAssemblyEnd(C, MAT_FINAL_ASSEMBLY));
102: PetscFunctionReturn(PETSC_SUCCESS);
103: }