Actual source code: mcomposite.c

  1: #include <../src/mat/impls/shell/shell.h>

  3: const char *const MatCompositeMergeTypes[] = {"left", "right", "MatCompositeMergeType", "MAT_COMPOSITE_", NULL};

  5: typedef struct _Mat_CompositeLink *Mat_CompositeLink;
  6: struct _Mat_CompositeLink {
  7:   Mat               mat;
  8:   Vec               work;
  9:   Mat_CompositeLink next, prev;
 10: };

 12: typedef struct {
 13:   MatCompositeType      type;
 14:   Mat_CompositeLink     head, tail;
 15:   Vec                   work;
 16:   PetscInt              nmat;
 17:   PetscBool             merge;
 18:   MatCompositeMergeType mergetype;
 19:   MatStructure          structure;

 21:   PetscScalar *scalings;
 22:   PetscBool    merge_mvctx; /* Whether need to merge mvctx of component matrices */
 23:   Vec         *lvecs;       /* [nmat] Basically, they are Mvctx->lvec of each component matrix */
 24:   PetscScalar *larray;      /* [len] Data arrays of lvecs[] are stored consecutively in larray */
 25:   PetscInt     len;         /* Length of larray[] */
 26:   Vec          gvec;        /* Union of lvecs[] without duplicated entries */
 27:   PetscInt    *location;    /* A map that maps entries in garray[] to larray[] */
 28:   VecScatter   Mvctx;
 29: } Mat_Composite;

 31: static PetscErrorCode MatCompositeDestroyMergedMvctx_Private(Mat_Composite *shell)
 32: {
 33:   PetscInt i;

 35:   PetscFunctionBegin;
 36:   if (shell->Mvctx) {
 37:     for (i = 0; i < shell->nmat; i++) PetscCall(VecDestroy(&shell->lvecs[i]));
 38:     PetscCall(PetscFree3(shell->location, shell->larray, shell->lvecs));
 39:     PetscCall(VecDestroy(&shell->gvec));
 40:     PetscCall(VecScatterDestroy(&shell->Mvctx));
 41:     shell->len = 0;
 42:   }
 43:   PetscFunctionReturn(PETSC_SUCCESS);
 44: }

 46: static PetscErrorCode MatDestroy_Composite(Mat mat)
 47: {
 48:   Mat_Composite    *shell;
 49:   Mat_CompositeLink next, oldnext;

 51:   PetscFunctionBegin;
 52:   PetscCall(MatShellGetContext(mat, &shell));
 53:   next = shell->head;
 54:   while (next) {
 55:     PetscCall(MatDestroy(&next->mat));
 56:     if (next->work && (!next->next || next->work != next->next->work)) PetscCall(VecDestroy(&next->work));
 57:     oldnext = next;
 58:     next    = next->next;
 59:     PetscCall(PetscFree(oldnext));
 60:   }
 61:   PetscCall(VecDestroy(&shell->work));

 63:   PetscCall(MatCompositeDestroyMergedMvctx_Private(shell));

 65:   PetscCall(PetscFree(shell->scalings));
 66:   PetscCall(PetscObjectComposeFunction((PetscObject)mat, "MatCompositeAddMat_C", NULL));
 67:   PetscCall(PetscObjectComposeFunction((PetscObject)mat, "MatCompositeSetType_C", NULL));
 68:   PetscCall(PetscObjectComposeFunction((PetscObject)mat, "MatCompositeGetType_C", NULL));
 69:   PetscCall(PetscObjectComposeFunction((PetscObject)mat, "MatCompositeSetMergeType_C", NULL));
 70:   PetscCall(PetscObjectComposeFunction((PetscObject)mat, "MatCompositeSetMatStructure_C", NULL));
 71:   PetscCall(PetscObjectComposeFunction((PetscObject)mat, "MatCompositeGetMatStructure_C", NULL));
 72:   PetscCall(PetscObjectComposeFunction((PetscObject)mat, "MatCompositeMerge_C", NULL));
 73:   PetscCall(PetscObjectComposeFunction((PetscObject)mat, "MatCompositeGetNumberMat_C", NULL));
 74:   PetscCall(PetscObjectComposeFunction((PetscObject)mat, "MatCompositeGetMat_C", NULL));
 75:   PetscCall(PetscObjectComposeFunction((PetscObject)mat, "MatCompositeSetScalings_C", NULL));
 76:   PetscCall(PetscFree(shell));
 77:   PetscCall(PetscObjectComposeFunction((PetscObject)mat, "MatShellSetContext_C", NULL)); // needed to avoid a call to MatShellSetContext_Immutable()
 78:   PetscFunctionReturn(PETSC_SUCCESS);
 79: }

 81: static PetscErrorCode MatMult_Composite_Multiplicative(Mat A, Vec x, Vec y)
 82: {
 83:   Mat_Composite    *shell;
 84:   Mat_CompositeLink next;
 85:   Vec               out;

 87:   PetscFunctionBegin;
 88:   PetscCall(MatShellGetContext(A, &shell));
 89:   next = shell->head;
 90:   PetscCheck(next, PETSC_COMM_SELF, PETSC_ERR_ARG_WRONGSTATE, "Must provide at least one matrix with MatCompositeAddMat()");
 91:   while (next->next) {
 92:     if (!next->work) { /* should reuse previous work if the same size */
 93:       PetscCall(MatCreateVecs(next->mat, NULL, &next->work));
 94:     }
 95:     out = next->work;
 96:     PetscCall(MatMult(next->mat, x, out));
 97:     x    = out;
 98:     next = next->next;
 99:   }
100:   PetscCall(MatMult(next->mat, x, y));
101:   if (shell->scalings) {
102:     PetscScalar scale = 1.0;
103:     for (PetscInt i = 0; i < shell->nmat; i++) scale *= shell->scalings[i];
104:     PetscCall(VecScale(y, scale));
105:   }
106:   PetscFunctionReturn(PETSC_SUCCESS);
107: }

109: static PetscErrorCode MatMultTranspose_Composite_Multiplicative(Mat A, Vec x, Vec y)
110: {
111:   Mat_Composite    *shell;
112:   Mat_CompositeLink tail;
113:   Vec               out;

115:   PetscFunctionBegin;
116:   PetscCall(MatShellGetContext(A, &shell));
117:   tail = shell->tail;
118:   PetscCheck(tail, PETSC_COMM_SELF, PETSC_ERR_ARG_WRONGSTATE, "Must provide at least one matrix with MatCompositeAddMat()");
119:   while (tail->prev) {
120:     if (!tail->prev->work) { /* should reuse previous work if the same size */
121:       PetscCall(MatCreateVecs(tail->mat, NULL, &tail->prev->work));
122:     }
123:     out = tail->prev->work;
124:     PetscCall(MatMultTranspose(tail->mat, x, out));
125:     x    = out;
126:     tail = tail->prev;
127:   }
128:   PetscCall(MatMultTranspose(tail->mat, x, y));
129:   if (shell->scalings) {
130:     PetscScalar scale = 1.0;
131:     for (PetscInt i = 0; i < shell->nmat; i++) scale *= shell->scalings[i];
132:     PetscCall(VecScale(y, scale));
133:   }
134:   PetscFunctionReturn(PETSC_SUCCESS);
135: }

137: static PetscErrorCode MatMult_Composite(Mat mat, Vec x, Vec y)
138: {
139:   Mat_Composite     *shell;
140:   Mat_CompositeLink  cur;
141:   Vec                y2, xin;
142:   Mat                A, B;
143:   PetscInt           i, j, k, n, nuniq, lo, hi, mid, *gindices, *buf, *tmp, tot;
144:   const PetscScalar *vals;
145:   const PetscInt    *garray;
146:   IS                 ix, iy;
147:   PetscBool          match;

149:   PetscFunctionBegin;
150:   PetscCall(MatShellGetContext(mat, &shell));
151:   cur = shell->head;
152:   PetscCheck(cur, PETSC_COMM_SELF, PETSC_ERR_ARG_WRONGSTATE, "Must provide at least one matrix with MatCompositeAddMat()");

154:   /* Try to merge Mvctx when instructed but not yet done. We did not do it in MatAssemblyEnd() since at that time
155:      we did not know whether mat is ADDITIVE or MULTIPLICATIVE. Only now we are assured mat is ADDITIVE and
156:      it is legal to merge Mvctx, because all component matrices have the same size.
157:    */
158:   if (shell->merge_mvctx && !shell->Mvctx) {
159:     /* Currently only implemented for MATMPIAIJ */
160:     for (cur = shell->head; cur; cur = cur->next) {
161:       PetscCall(PetscObjectTypeCompare((PetscObject)cur->mat, MATMPIAIJ, &match));
162:       if (!match) {
163:         shell->merge_mvctx = PETSC_FALSE;
164:         goto skip_merge_mvctx;
165:       }
166:     }

168:     /* Go through matrices first time to count total number of nonzero off-diag columns (may have dups) */
169:     tot = 0;
170:     for (cur = shell->head; cur; cur = cur->next) {
171:       PetscCall(MatMPIAIJGetSeqAIJ(cur->mat, NULL, &B, NULL));
172:       PetscCall(MatGetLocalSize(B, NULL, &n));
173:       tot += n;
174:     }
175:     PetscCall(PetscMalloc3(tot, &shell->location, tot, &shell->larray, shell->nmat, &shell->lvecs));
176:     shell->len = tot;

178:     /* Go through matrices second time to sort off-diag columns and remove dups */
179:     PetscCall(PetscMalloc1(tot, &gindices)); /* No Malloc2() since we will give one to PETSc and free the other */
180:     PetscCall(PetscMalloc1(tot, &buf));
181:     nuniq = 0; /* Number of unique nonzero columns */
182:     for (cur = shell->head; cur; cur = cur->next) {
183:       PetscCall(MatMPIAIJGetSeqAIJ(cur->mat, NULL, &B, &garray));
184:       PetscCall(MatGetLocalSize(B, NULL, &n));
185:       /* Merge pre-sorted garray[0,n) and gindices[0,nuniq) to buf[] */
186:       i = j = k = 0;
187:       while (i < n && j < nuniq) {
188:         if (garray[i] < gindices[j]) buf[k++] = garray[i++];
189:         else if (garray[i] > gindices[j]) buf[k++] = gindices[j++];
190:         else {
191:           buf[k++] = garray[i++];
192:           j++;
193:         }
194:       }
195:       /* Copy leftover in garray[] or gindices[] */
196:       if (i < n) {
197:         PetscCall(PetscArraycpy(buf + k, garray + i, n - i));
198:         nuniq = k + n - i;
199:       } else if (j < nuniq) {
200:         PetscCall(PetscArraycpy(buf + k, gindices + j, nuniq - j));
201:         nuniq = k + nuniq - j;
202:       } else nuniq = k;
203:       /* Swap gindices and buf to merge garray of the next matrix */
204:       tmp      = gindices;
205:       gindices = buf;
206:       buf      = tmp;
207:     }
208:     PetscCall(PetscFree(buf));

210:     /* Go through matrices third time to build a map from gindices[] to garray[] */
211:     tot = 0;
212:     for (cur = shell->head, j = 0; cur; cur = cur->next, j++) { /* j-th matrix */
213:       PetscCall(MatMPIAIJGetSeqAIJ(cur->mat, NULL, &B, &garray));
214:       PetscCall(MatGetLocalSize(B, NULL, &n));
215:       PetscCall(VecCreateSeqWithArray(PETSC_COMM_SELF, 1, n, NULL, &shell->lvecs[j]));
216:       /* This is an optimized PetscFindInt(garray[i],nuniq,gindices,&shell->location[tot+i]), using the fact that garray[] is also sorted */
217:       lo = 0;
218:       for (i = 0; i < n; i++) {
219:         hi = nuniq;
220:         while (hi - lo > 1) {
221:           mid = lo + (hi - lo) / 2;
222:           if (garray[i] < gindices[mid]) hi = mid;
223:           else lo = mid;
224:         }
225:         shell->location[tot + i] = lo; /* gindices[lo] = garray[i] */
226:         lo++;                          /* Since garray[i+1] > garray[i], we can safely advance lo */
227:       }
228:       tot += n;
229:     }

231:     /* Build merged Mvctx */
232:     PetscCall(ISCreateGeneral(PETSC_COMM_SELF, nuniq, gindices, PETSC_OWN_POINTER, &ix));
233:     PetscCall(ISCreateStride(PETSC_COMM_SELF, nuniq, 0, 1, &iy));
234:     PetscCall(VecCreateMPIWithArray(PetscObjectComm((PetscObject)mat), 1, mat->cmap->n, mat->cmap->N, NULL, &xin));
235:     PetscCall(VecCreateSeq(PETSC_COMM_SELF, nuniq, &shell->gvec));
236:     PetscCall(VecScatterCreate(xin, ix, shell->gvec, iy, &shell->Mvctx));
237:     PetscCall(VecDestroy(&xin));
238:     PetscCall(ISDestroy(&ix));
239:     PetscCall(ISDestroy(&iy));
240:   }

242: skip_merge_mvctx:
243:   PetscCall(VecSet(y, 0));
244:   if (!((Mat_Shell *)mat->data)->left_work) PetscCall(VecDuplicate(y, &(((Mat_Shell *)mat->data)->left_work)));
245:   y2 = ((Mat_Shell *)mat->data)->left_work;

247:   if (shell->Mvctx) { /* Have a merged Mvctx */
248:     /* Suppose we want to compute y = sMx, where s is the scaling factor and A, B are matrix M's diagonal/off-diagonal part. We could do
249:        in y = s(Ax1 + Bx2) or y = sAx1 + sBx2. The former incurs less FLOPS than the latter, but the latter provides an opportunity to
250:        overlap communication/computation since we can do sAx1 while communicating x2. Here, we use the former approach.
251:      */
252:     PetscCall(VecScatterBegin(shell->Mvctx, x, shell->gvec, INSERT_VALUES, SCATTER_FORWARD));
253:     PetscCall(VecScatterEnd(shell->Mvctx, x, shell->gvec, INSERT_VALUES, SCATTER_FORWARD));

255:     PetscCall(VecGetArrayRead(shell->gvec, &vals));
256:     for (i = 0; i < shell->len; i++) shell->larray[i] = vals[shell->location[i]];
257:     PetscCall(VecRestoreArrayRead(shell->gvec, &vals));

259:     for (cur = shell->head, tot = i = 0; cur; cur = cur->next, i++) { /* i-th matrix */
260:       PetscCall(MatMPIAIJGetSeqAIJ(cur->mat, &A, &B, NULL));
261:       PetscUseTypeMethod(A, mult, x, y2);
262:       PetscCall(MatGetLocalSize(B, NULL, &n));
263:       PetscCall(VecPlaceArray(shell->lvecs[i], &shell->larray[tot]));
264:       PetscUseTypeMethod(B, multadd, shell->lvecs[i], y2, y2);
265:       PetscCall(VecResetArray(shell->lvecs[i]));
266:       PetscCall(VecAXPY(y, shell->scalings ? shell->scalings[i] : 1.0, y2));
267:       tot += n;
268:     }
269:   } else {
270:     if (shell->scalings) {
271:       for (cur = shell->head, i = 0; cur; cur = cur->next, i++) {
272:         PetscCall(MatMult(cur->mat, x, y2));
273:         PetscCall(VecAXPY(y, shell->scalings[i], y2));
274:       }
275:     } else {
276:       for (cur = shell->head; cur; cur = cur->next) PetscCall(MatMultAdd(cur->mat, x, y, y));
277:     }
278:   }
279:   PetscFunctionReturn(PETSC_SUCCESS);
280: }

282: static PetscErrorCode MatMultTranspose_Composite(Mat A, Vec x, Vec y)
283: {
284:   Mat_Composite    *shell;
285:   Mat_CompositeLink next;
286:   Vec               y2 = NULL;
287:   PetscInt          i;

289:   PetscFunctionBegin;
290:   PetscCall(MatShellGetContext(A, &shell));
291:   next = shell->head;
292:   PetscCheck(next, PETSC_COMM_SELF, PETSC_ERR_ARG_WRONGSTATE, "Must provide at least one matrix with MatCompositeAddMat()");

294:   PetscCall(MatMultTranspose(next->mat, x, y));
295:   if (shell->scalings) {
296:     PetscCall(VecScale(y, shell->scalings[0]));
297:     if (!((Mat_Shell *)A->data)->right_work) PetscCall(VecDuplicate(y, &(((Mat_Shell *)A->data)->right_work)));
298:     y2 = ((Mat_Shell *)A->data)->right_work;
299:   }
300:   i = 1;
301:   while ((next = next->next)) {
302:     if (!shell->scalings) PetscCall(MatMultTransposeAdd(next->mat, x, y, y));
303:     else {
304:       PetscCall(MatMultTranspose(next->mat, x, y2));
305:       PetscCall(VecAXPY(y, shell->scalings[i++], y2));
306:     }
307:   }
308:   PetscFunctionReturn(PETSC_SUCCESS);
309: }

311: static PetscErrorCode MatGetDiagonal_Composite(Mat A, Vec v)
312: {
313:   Mat_Composite    *shell;
314:   Mat_CompositeLink next;
315:   PetscInt          i;

317:   PetscFunctionBegin;
318:   PetscCall(MatShellGetContext(A, &shell));
319:   next = shell->head;
320:   PetscCheck(next, PETSC_COMM_SELF, PETSC_ERR_ARG_WRONGSTATE, "Must provide at least one matrix with MatCompositeAddMat()");
321:   PetscCall(MatGetDiagonal(next->mat, v));
322:   if (shell->scalings) PetscCall(VecScale(v, shell->scalings[0]));

324:   if (next->next && !shell->work) PetscCall(VecDuplicate(v, &shell->work));
325:   i = 1;
326:   while ((next = next->next)) {
327:     PetscCall(MatGetDiagonal(next->mat, shell->work));
328:     PetscCall(VecAXPY(v, shell->scalings ? shell->scalings[i++] : 1.0, shell->work));
329:   }
330:   PetscFunctionReturn(PETSC_SUCCESS);
331: }

333: static PetscErrorCode MatAssemblyEnd_Composite(Mat Y, MatAssemblyType t)
334: {
335:   Mat_Composite *shell;

337:   PetscFunctionBegin;
338:   PetscCall(MatShellGetContext(Y, &shell));
339:   if (shell->merge) PetscCall(MatCompositeMerge(Y));
340:   else PetscCall(MatAssemblyEnd_Shell(Y, t));
341:   PetscFunctionReturn(PETSC_SUCCESS);
342: }

344: static PetscErrorCode MatSetFromOptions_Composite(Mat A, PetscOptionItems PetscOptionsObject)
345: {
346:   Mat_Composite *a;

348:   PetscFunctionBegin;
349:   PetscCall(MatShellGetContext(A, &a));
350:   PetscOptionsHeadBegin(PetscOptionsObject, "MATCOMPOSITE options");
351:   PetscCall(PetscOptionsBool("-mat_composite_merge", "Merge at MatAssemblyEnd", "MatCompositeMerge", a->merge, &a->merge, NULL));
352:   PetscCall(PetscOptionsEnum("-mat_composite_merge_type", "Set composite merge direction", "MatCompositeSetMergeType", MatCompositeMergeTypes, (PetscEnum)a->mergetype, (PetscEnum *)&a->mergetype, NULL));
353:   PetscCall(PetscOptionsBool("-mat_composite_merge_mvctx", "Merge MatMult() vecscat contexts", "MatCreateComposite", a->merge_mvctx, &a->merge_mvctx, NULL));
354:   PetscOptionsHeadEnd();
355:   PetscFunctionReturn(PETSC_SUCCESS);
356: }

358: /*@
359:   MatCreateComposite - Creates a matrix as the sum or product of one or more matrices

361:   Collective

363:   Input Parameters:
364: + comm - MPI communicator
365: . nmat - number of matrices to put in
366: - mats - the matrices

368:   Output Parameter:
369: . mat - the matrix

371:   Options Database Keys:
372: + -mat_composite_merge       - merge in `MatAssemblyEnd()`
373: . -mat_composite_merge_mvctx - merge Mvctx of component matrices to optimize communication in `MatMult()` for ADDITIVE matrices
374: - -mat_composite_merge_type  - set merge direction

376:   Level: advanced

378:   Note:
379:   Alternative construction
380: .vb
381:        MatCreate(comm,&mat);
382:        MatSetSizes(mat,m,n,M,N);
383:        MatSetType(mat,MATCOMPOSITE);
384:        MatCompositeAddMat(mat,mats[0]);
385:        ....
386:        MatCompositeAddMat(mat,mats[nmat-1]);
387:        MatAssemblyBegin(mat,MAT_FINAL_ASSEMBLY);
388:        MatAssemblyEnd(mat,MAT_FINAL_ASSEMBLY);
389: .ve

391:   For the multiplicative form the product is mat[nmat-1]*mat[nmat-2]*....*mat[0]

393: .seealso: [](ch_matrices), `Mat`, `MatDestroy()`, `MatMult()`, `MatCompositeAddMat()`, `MatCompositeGetMat()`, `MatCompositeMerge()`, `MatCompositeSetType()`,
394:           `MATCOMPOSITE`, `MatCompositeType`
395: @*/
396: PetscErrorCode MatCreateComposite(MPI_Comm comm, PetscInt nmat, const Mat *mats, Mat *mat)
397: {
398:   PetscFunctionBegin;
399:   PetscCheck(nmat >= 1, PETSC_COMM_SELF, PETSC_ERR_ARG_OUTOFRANGE, "Must pass in at least one matrix");
400:   PetscAssertPointer(mat, 4);
401:   PetscCall(MatCreate(comm, mat));
402:   PetscCall(MatSetType(*mat, MATCOMPOSITE));
403:   for (PetscInt i = 0; i < nmat; i++) PetscCall(MatCompositeAddMat(*mat, mats[i]));
404:   PetscCall(MatAssemblyBegin(*mat, MAT_FINAL_ASSEMBLY));
405:   PetscCall(MatAssemblyEnd(*mat, MAT_FINAL_ASSEMBLY));
406:   PetscFunctionReturn(PETSC_SUCCESS);
407: }

409: static PetscErrorCode MatCompositeAddMat_Composite(Mat mat, Mat smat)
410: {
411:   Mat_Composite    *shell;
412:   Mat_CompositeLink ilink, next;
413:   VecType           vtype_mat, vtype_smat;
414:   PetscBool         match;

416:   PetscFunctionBegin;
417:   PetscCall(MatShellGetContext(mat, &shell));
418:   PetscCall(MatCompositeDestroyMergedMvctx_Private(shell));
419:   next = shell->head;
420:   PetscCall(PetscNew(&ilink));
421:   ilink->next = NULL;
422:   PetscCall(PetscObjectReference((PetscObject)smat));
423:   ilink->mat = smat;

425:   if (!next) shell->head = ilink;
426:   else {
427:     while (next->next) next = next->next;
428:     next->next  = ilink;
429:     ilink->prev = next;
430:   }
431:   shell->tail = ilink;
432:   shell->nmat += 1;

434:   /* If all of the partial matrices have the same default vector type, then the composite matrix should also have this default type.
435:      Otherwise, the default type should be "standard". */
436:   PetscCall(MatGetVecType(smat, &vtype_smat));
437:   if (shell->nmat == 1) PetscCall(MatSetVecType(mat, vtype_smat));
438:   else {
439:     PetscCall(MatGetVecType(mat, &vtype_mat));
440:     PetscCall(PetscStrcmp(vtype_smat, vtype_mat, &match));
441:     if (!match) PetscCall(MatSetVecType(mat, VECSTANDARD));
442:   }

444:   /* Retain the old scalings (if any) and expand it with a 1.0 for the newly added matrix */
445:   if (shell->scalings) {
446:     PetscCall(PetscRealloc(sizeof(PetscScalar) * shell->nmat, &shell->scalings));
447:     shell->scalings[shell->nmat - 1] = 1.0;
448:   }

450:   /* The composite matrix requires PetscLayouts for its rows and columns; we copy these from the constituent partial matrices. */
451:   if (shell->nmat == 1) PetscCall(PetscLayoutReference(smat->cmap, &mat->cmap));
452:   PetscCall(PetscLayoutReference(smat->rmap, &mat->rmap));
453:   PetscFunctionReturn(PETSC_SUCCESS);
454: }

456: /*@
457:   MatCompositeAddMat - Add another matrix to a composite matrix.

459:   Collective

461:   Input Parameters:
462: + mat  - the composite matrix
463: - smat - the partial matrix

465:   Level: advanced

467: .seealso: [](ch_matrices), `Mat`, `MatCreateComposite()`, `MatCompositeGetMat()`, `MATCOMPOSITE`
468: @*/
469: PetscErrorCode MatCompositeAddMat(Mat mat, Mat smat)
470: {
471:   PetscFunctionBegin;
474:   PetscUseMethod(mat, "MatCompositeAddMat_C", (Mat, Mat), (mat, smat));
475:   PetscFunctionReturn(PETSC_SUCCESS);
476: }

478: static PetscErrorCode MatCompositeSetType_Composite(Mat mat, MatCompositeType type)
479: {
480:   Mat_Composite *b;

482:   PetscFunctionBegin;
483:   PetscCall(MatShellGetContext(mat, &b));
484:   b->type = type;
485:   if (type == MAT_COMPOSITE_MULTIPLICATIVE) {
486:     PetscCall(MatShellSetOperation(mat, MATOP_GET_DIAGONAL, NULL));
487:     PetscCall(MatShellSetOperation(mat, MATOP_MULT, (PetscErrorCodeFn *)MatMult_Composite_Multiplicative));
488:     PetscCall(MatShellSetOperation(mat, MATOP_MULT_TRANSPOSE, (PetscErrorCodeFn *)MatMultTranspose_Composite_Multiplicative));
489:     b->merge_mvctx = PETSC_FALSE;
490:   } else {
491:     PetscCall(MatShellSetOperation(mat, MATOP_GET_DIAGONAL, (PetscErrorCodeFn *)MatGetDiagonal_Composite));
492:     PetscCall(MatShellSetOperation(mat, MATOP_MULT, (PetscErrorCodeFn *)MatMult_Composite));
493:     PetscCall(MatShellSetOperation(mat, MATOP_MULT_TRANSPOSE, (PetscErrorCodeFn *)MatMultTranspose_Composite));
494:   }
495:   PetscFunctionReturn(PETSC_SUCCESS);
496: }

498: /*@
499:   MatCompositeSetType - Indicates if the matrix is defined as the sum of a set of matrices or the product.

501:   Logically Collective

503:   Input Parameters:
504: + mat  - the composite matrix
505: - type - the `MatCompositeType` to use for the matrix

507:   Level: advanced

509: .seealso: [](ch_matrices), `Mat`, `MatDestroy()`, `MatMult()`, `MatCompositeAddMat()`, `MatCreateComposite()`, `MatCompositeGetType()`, `MATCOMPOSITE`,
510:           `MatCompositeType`
511: @*/
512: PetscErrorCode MatCompositeSetType(Mat mat, MatCompositeType type)
513: {
514:   PetscFunctionBegin;
517:   PetscUseMethod(mat, "MatCompositeSetType_C", (Mat, MatCompositeType), (mat, type));
518:   PetscFunctionReturn(PETSC_SUCCESS);
519: }

521: static PetscErrorCode MatCompositeGetType_Composite(Mat mat, MatCompositeType *type)
522: {
523:   Mat_Composite *shell;

525:   PetscFunctionBegin;
526:   PetscCall(MatShellGetContext(mat, &shell));
527:   *type = shell->type;
528:   PetscFunctionReturn(PETSC_SUCCESS);
529: }

531: /*@
532:   MatCompositeGetType - Returns type of composite.

534:   Not Collective

536:   Input Parameter:
537: . mat - the composite matrix

539:   Output Parameter:
540: . type - type of composite

542:   Level: advanced

544: .seealso: [](ch_matrices), `Mat`, `MatCreateComposite()`, `MatCompositeSetType()`, `MATCOMPOSITE`, `MatCompositeType`
545: @*/
546: PetscErrorCode MatCompositeGetType(Mat mat, MatCompositeType *type)
547: {
548:   PetscFunctionBegin;
550:   PetscAssertPointer(type, 2);
551:   PetscUseMethod(mat, "MatCompositeGetType_C", (Mat, MatCompositeType *), (mat, type));
552:   PetscFunctionReturn(PETSC_SUCCESS);
553: }

555: static PetscErrorCode MatCompositeSetMatStructure_Composite(Mat mat, MatStructure str)
556: {
557:   Mat_Composite *shell;

559:   PetscFunctionBegin;
560:   PetscCall(MatShellGetContext(mat, &shell));
561:   shell->structure = str;
562:   PetscFunctionReturn(PETSC_SUCCESS);
563: }

565: /*@
566:   MatCompositeSetMatStructure - Indicates structure of matrices in the composite matrix.

568:   Not Collective

570:   Input Parameters:
571: + mat - the composite matrix
572: - str - either `SAME_NONZERO_PATTERN`, `DIFFERENT_NONZERO_PATTERN` (default) or `SUBSET_NONZERO_PATTERN`

574:   Level: advanced

576:   Note:
577:   Information about the matrices structure is used in `MatCompositeMerge()` for additive composite matrix.

579: .seealso: [](ch_matrices), `Mat`, `MatAXPY()`, `MatCreateComposite()`, `MatCompositeMerge()`, `MatCompositeGetMatStructure()`, `MATCOMPOSITE`
580: @*/
581: PetscErrorCode MatCompositeSetMatStructure(Mat mat, MatStructure str)
582: {
583:   PetscFunctionBegin;
585:   PetscUseMethod(mat, "MatCompositeSetMatStructure_C", (Mat, MatStructure), (mat, str));
586:   PetscFunctionReturn(PETSC_SUCCESS);
587: }

589: static PetscErrorCode MatCompositeGetMatStructure_Composite(Mat mat, MatStructure *str)
590: {
591:   Mat_Composite *shell;

593:   PetscFunctionBegin;
594:   PetscCall(MatShellGetContext(mat, &shell));
595:   *str = shell->structure;
596:   PetscFunctionReturn(PETSC_SUCCESS);
597: }

599: /*@
600:   MatCompositeGetMatStructure - Returns the structure of matrices in the composite matrix.

602:   Not Collective

604:   Input Parameter:
605: . mat - the composite matrix

607:   Output Parameter:
608: . str - structure of the matrices

610:   Level: advanced

612: .seealso: [](ch_matrices), `Mat`, `MatCreateComposite()`, `MatCompositeSetMatStructure()`, `MATCOMPOSITE`
613: @*/
614: PetscErrorCode MatCompositeGetMatStructure(Mat mat, MatStructure *str)
615: {
616:   PetscFunctionBegin;
618:   PetscAssertPointer(str, 2);
619:   PetscUseMethod(mat, "MatCompositeGetMatStructure_C", (Mat, MatStructure *), (mat, str));
620:   PetscFunctionReturn(PETSC_SUCCESS);
621: }

623: static PetscErrorCode MatCompositeSetMergeType_Composite(Mat mat, MatCompositeMergeType type)
624: {
625:   Mat_Composite *shell;

627:   PetscFunctionBegin;
628:   PetscCall(MatShellGetContext(mat, &shell));
629:   shell->mergetype = type;
630:   PetscFunctionReturn(PETSC_SUCCESS);
631: }

633: /*@
634:   MatCompositeSetMergeType - Sets order of `MatCompositeMerge()`.

636:   Logically Collective

638:   Input Parameters:
639: + mat  - the composite matrix
640: - type - `MAT_COMPOSITE_MERGE RIGHT` (default) to start merge from right with the first added matrix (mat[0]),
641:           `MAT_COMPOSITE_MERGE_LEFT` to start merge from left with the last added matrix (mat[nmat-1])

643:   Level: advanced

645:   Note:
646:   The resulting matrix is the same regardless of the `MatCompositeMergeType`. Only the order of operation is changed.
647:   If set to `MAT_COMPOSITE_MERGE_RIGHT` the order of the merge is mat[nmat-1]*(mat[nmat-2]*(...*(mat[1]*mat[0])))
648:   otherwise the order is (((mat[nmat-1]*mat[nmat-2])*mat[nmat-3])*...)*mat[0].

650: .seealso: [](ch_matrices), `Mat`, `MatCreateComposite()`, `MatCompositeMerge()`, `MATCOMPOSITE`
651: @*/
652: PetscErrorCode MatCompositeSetMergeType(Mat mat, MatCompositeMergeType type)
653: {
654:   PetscFunctionBegin;
657:   PetscUseMethod(mat, "MatCompositeSetMergeType_C", (Mat, MatCompositeMergeType), (mat, type));
658:   PetscFunctionReturn(PETSC_SUCCESS);
659: }

661: static PetscErrorCode MatCompositeMerge_Composite(Mat mat)
662: {
663:   Mat_Composite    *shell;
664:   Mat_CompositeLink next, prev;
665:   Mat               tmat, newmat;
666:   Vec               left, right, dshift;
667:   PetscScalar       scale, shift;
668:   PetscInt          i;

670:   PetscFunctionBegin;
671:   PetscCall(MatShellGetContext(mat, &shell));
672:   next = shell->head;
673:   prev = shell->tail;
674:   PetscCheck(next, PETSC_COMM_SELF, PETSC_ERR_ARG_WRONGSTATE, "Must provide at least one matrix with MatCompositeAddMat()");
675:   PetscCall(MatShellGetScalingShifts(mat, &shift, &scale, &dshift, &left, &right, (Mat *)MAT_SHELL_NOT_ALLOWED, (IS *)MAT_SHELL_NOT_ALLOWED, (IS *)MAT_SHELL_NOT_ALLOWED));
676:   if (shell->type == MAT_COMPOSITE_ADDITIVE) {
677:     if (shell->mergetype == MAT_COMPOSITE_MERGE_RIGHT) {
678:       i = 0;
679:       PetscCall(MatDuplicate(next->mat, MAT_COPY_VALUES, &tmat));
680:       if (shell->scalings) PetscCall(MatScale(tmat, shell->scalings[i++]));
681:       while ((next = next->next)) PetscCall(MatAXPY(tmat, shell->scalings ? shell->scalings[i++] : 1.0, next->mat, shell->structure));
682:     } else {
683:       i = shell->nmat - 1;
684:       PetscCall(MatDuplicate(prev->mat, MAT_COPY_VALUES, &tmat));
685:       if (shell->scalings) PetscCall(MatScale(tmat, shell->scalings[i--]));
686:       while ((prev = prev->prev)) PetscCall(MatAXPY(tmat, shell->scalings ? shell->scalings[i--] : 1.0, prev->mat, shell->structure));
687:     }
688:   } else {
689:     if (shell->mergetype == MAT_COMPOSITE_MERGE_RIGHT) {
690:       PetscCall(MatDuplicate(next->mat, MAT_COPY_VALUES, &tmat));
691:       while ((next = next->next)) {
692:         PetscCall(MatMatMult(next->mat, tmat, MAT_INITIAL_MATRIX, PETSC_DETERMINE, &newmat));
693:         PetscCall(MatDestroy(&tmat));
694:         tmat = newmat;
695:       }
696:     } else {
697:       PetscCall(MatDuplicate(prev->mat, MAT_COPY_VALUES, &tmat));
698:       while ((prev = prev->prev)) {
699:         PetscCall(MatMatMult(tmat, prev->mat, MAT_INITIAL_MATRIX, PETSC_DETERMINE, &newmat));
700:         PetscCall(MatDestroy(&tmat));
701:         tmat = newmat;
702:       }
703:     }
704:     if (shell->scalings) {
705:       for (i = 0; i < shell->nmat; i++) scale *= shell->scalings[i];
706:     }
707:   }

709:   PetscCall(PetscObjectReference((PetscObject)left));
710:   PetscCall(PetscObjectReference((PetscObject)right));
711:   PetscCall(PetscObjectReference((PetscObject)dshift));

713:   PetscCall(MatHeaderReplace(mat, &tmat));

715:   PetscCall(MatDiagonalScale(mat, left, right));
716:   PetscCall(MatScale(mat, scale));
717:   PetscCall(MatShift(mat, shift));
718:   PetscCall(VecDestroy(&left));
719:   PetscCall(VecDestroy(&right));
720:   if (dshift) {
721:     PetscCall(MatDiagonalSet(mat, dshift, ADD_VALUES));
722:     PetscCall(VecDestroy(&dshift));
723:   }
724:   PetscFunctionReturn(PETSC_SUCCESS);
725: }

727: /*@
728:   MatCompositeMerge - Given a composite matrix, replaces it with a "regular" matrix
729:   by summing or computing the product of all the matrices inside the composite matrix.

731:   Collective

733:   Input Parameter:
734: . mat - the composite matrix

736:   Options Database Keys:
737: + -mat_composite_merge      - merge in `MatAssemblyEnd()`
738: - -mat_composite_merge_type - set merge direction

740:   Level: advanced

742:   Note:
743:   The `MatType` of the resulting matrix will be the same as the `MatType` of the FIRST matrix in the composite matrix.

745: .seealso: [](ch_matrices), `Mat`, `MatDestroy()`, `MatMult()`, `MatCompositeAddMat()`, `MatCreateComposite()`, `MatCompositeSetMatStructure()`, `MatCompositeSetMergeType()`, `MATCOMPOSITE`
746: @*/
747: PetscErrorCode MatCompositeMerge(Mat mat)
748: {
749:   PetscFunctionBegin;
751:   PetscUseMethod(mat, "MatCompositeMerge_C", (Mat), (mat));
752:   PetscFunctionReturn(PETSC_SUCCESS);
753: }

755: static PetscErrorCode MatCompositeGetNumberMat_Composite(Mat mat, PetscInt *nmat)
756: {
757:   Mat_Composite *shell;

759:   PetscFunctionBegin;
760:   PetscCall(MatShellGetContext(mat, &shell));
761:   *nmat = shell->nmat;
762:   PetscFunctionReturn(PETSC_SUCCESS);
763: }

765: /*@
766:   MatCompositeGetNumberMat - Returns the number of matrices in the composite matrix.

768:   Not Collective

770:   Input Parameter:
771: . mat - the composite matrix

773:   Output Parameter:
774: . nmat - number of matrices in the composite matrix

776:   Level: advanced

778: .seealso: [](ch_matrices), `Mat`, `MatCreateComposite()`, `MatCompositeGetMat()`, `MATCOMPOSITE`
779: @*/
780: PetscErrorCode MatCompositeGetNumberMat(Mat mat, PetscInt *nmat)
781: {
782:   PetscFunctionBegin;
784:   PetscAssertPointer(nmat, 2);
785:   PetscUseMethod(mat, "MatCompositeGetNumberMat_C", (Mat, PetscInt *), (mat, nmat));
786:   PetscFunctionReturn(PETSC_SUCCESS);
787: }

789: static PetscErrorCode MatCompositeGetMat_Composite(Mat mat, PetscInt i, Mat *Ai)
790: {
791:   Mat_Composite    *shell;
792:   Mat_CompositeLink ilink;

794:   PetscFunctionBegin;
795:   PetscCall(MatShellGetContext(mat, &shell));
796:   PetscCheck(i < shell->nmat, PetscObjectComm((PetscObject)mat), PETSC_ERR_ARG_OUTOFRANGE, "index out of range: %" PetscInt_FMT " >= %" PetscInt_FMT, i, shell->nmat);
797:   ilink = shell->head;
798:   for (PetscInt k = 0; k < i; k++) ilink = ilink->next;
799:   *Ai = ilink->mat;
800:   PetscFunctionReturn(PETSC_SUCCESS);
801: }

803: /*@
804:   MatCompositeGetMat - Returns the ith matrix from the composite matrix.

806:   Logically Collective

808:   Input Parameters:
809: + mat - the composite matrix
810: - i   - the number of requested matrix

812:   Output Parameter:
813: . Ai - ith matrix in composite

815:   Level: advanced

817: .seealso: [](ch_matrices), `Mat`, `MatCreateComposite()`, `MatCompositeGetNumberMat()`, `MatCompositeAddMat()`, `MATCOMPOSITE`
818: @*/
819: PetscErrorCode MatCompositeGetMat(Mat mat, PetscInt i, Mat *Ai)
820: {
821:   PetscFunctionBegin;
824:   PetscAssertPointer(Ai, 3);
825:   PetscUseMethod(mat, "MatCompositeGetMat_C", (Mat, PetscInt, Mat *), (mat, i, Ai));
826:   PetscFunctionReturn(PETSC_SUCCESS);
827: }

829: static PetscErrorCode MatCompositeSetScalings_Composite(Mat mat, const PetscScalar *scalings)
830: {
831:   Mat_Composite *shell;
832:   PetscInt       nmat;

834:   PetscFunctionBegin;
835:   PetscCall(MatShellGetContext(mat, &shell));
836:   PetscCall(MatCompositeGetNumberMat(mat, &nmat));
837:   if (!shell->scalings) PetscCall(PetscMalloc1(nmat, &shell->scalings));
838:   PetscCall(PetscArraycpy(shell->scalings, scalings, nmat));
839:   PetscFunctionReturn(PETSC_SUCCESS);
840: }

842: /*@
843:   MatCompositeSetScalings - Sets separate scaling factors for component matrices.

845:   Logically Collective

847:   Input Parameters:
848: + mat      - the composite matrix
849: - scalings - array of scaling factors with scalings[i] being factor of i-th matrix, for i in [0, nmat)

851:   Level: advanced

853: .seealso: [](ch_matrices), `Mat`, `MatScale()`, `MatDiagonalScale()`, `MATCOMPOSITE`
854: @*/
855: PetscErrorCode MatCompositeSetScalings(Mat mat, const PetscScalar *scalings)
856: {
857:   PetscFunctionBegin;
859:   PetscAssertPointer(scalings, 2);
861:   PetscUseMethod(mat, "MatCompositeSetScalings_C", (Mat, const PetscScalar *), (mat, scalings));
862:   PetscFunctionReturn(PETSC_SUCCESS);
863: }

865: /*MC
866:    MATCOMPOSITE - A matrix defined by the sum (or product) of one or more matrices.
867:     The matrices need to have a correct size and parallel layout for the sum or product to be valid.

869:   Level: advanced

871:    Note:
872:    To use the product of the matrices call `MatCompositeSetType`(mat,`MAT_COMPOSITE_MULTIPLICATIVE`);

874:   Developer Notes:
875:   This is implemented on top of `MATSHELL` to get support for scaling and shifting without requiring duplicate code

877:   Users can not call `MatShellSetOperation()` operations on this class, there is some error checking for that incorrect usage

879: .seealso: [](ch_matrices), `Mat`, `MatCreateComposite()`, `MatCompositeSetScalings()`, `MatCompositeAddMat()`, `MatSetType()`, `MatCompositeSetType()`, `MatCompositeGetType()`,
880:           `MatCompositeSetMatStructure()`, `MatCompositeGetMatStructure()`, `MatCompositeMerge()`, `MatCompositeSetMergeType()`, `MatCompositeGetNumberMat()`, `MatCompositeGetMat()`
881: M*/

883: PETSC_EXTERN PetscErrorCode MatCreate_Composite(Mat A)
884: {
885:   Mat_Composite *b;

887:   PetscFunctionBegin;
888:   PetscCall(PetscNew(&b));

890:   b->type        = MAT_COMPOSITE_ADDITIVE;
891:   b->nmat        = 0;
892:   b->merge       = PETSC_FALSE;
893:   b->mergetype   = MAT_COMPOSITE_MERGE_RIGHT;
894:   b->structure   = DIFFERENT_NONZERO_PATTERN;
895:   b->merge_mvctx = PETSC_TRUE;

897:   PetscCall(MatSetType(A, MATSHELL));
898:   PetscCall(MatShellSetContext(A, b));
899:   PetscCall(MatShellSetOperation(A, MATOP_DESTROY, (PetscErrorCodeFn *)MatDestroy_Composite));
900:   PetscCall(MatShellSetOperation(A, MATOP_MULT, (PetscErrorCodeFn *)MatMult_Composite));
901:   PetscCall(MatShellSetOperation(A, MATOP_MULT_TRANSPOSE, (PetscErrorCodeFn *)MatMultTranspose_Composite));
902:   PetscCall(MatShellSetOperation(A, MATOP_GET_DIAGONAL, (PetscErrorCodeFn *)MatGetDiagonal_Composite));
903:   PetscCall(MatShellSetOperation(A, MATOP_ASSEMBLY_END, (PetscErrorCodeFn *)MatAssemblyEnd_Composite));
904:   PetscCall(MatShellSetOperation(A, MATOP_SET_FROM_OPTIONS, (PetscErrorCodeFn *)MatSetFromOptions_Composite));
905:   PetscCall(PetscObjectComposeFunction((PetscObject)A, "MatCompositeAddMat_C", MatCompositeAddMat_Composite));
906:   PetscCall(PetscObjectComposeFunction((PetscObject)A, "MatCompositeSetType_C", MatCompositeSetType_Composite));
907:   PetscCall(PetscObjectComposeFunction((PetscObject)A, "MatCompositeGetType_C", MatCompositeGetType_Composite));
908:   PetscCall(PetscObjectComposeFunction((PetscObject)A, "MatCompositeSetMergeType_C", MatCompositeSetMergeType_Composite));
909:   PetscCall(PetscObjectComposeFunction((PetscObject)A, "MatCompositeSetMatStructure_C", MatCompositeSetMatStructure_Composite));
910:   PetscCall(PetscObjectComposeFunction((PetscObject)A, "MatCompositeGetMatStructure_C", MatCompositeGetMatStructure_Composite));
911:   PetscCall(PetscObjectComposeFunction((PetscObject)A, "MatCompositeMerge_C", MatCompositeMerge_Composite));
912:   PetscCall(PetscObjectComposeFunction((PetscObject)A, "MatCompositeGetNumberMat_C", MatCompositeGetNumberMat_Composite));
913:   PetscCall(PetscObjectComposeFunction((PetscObject)A, "MatCompositeGetMat_C", MatCompositeGetMat_Composite));
914:   PetscCall(PetscObjectComposeFunction((PetscObject)A, "MatCompositeSetScalings_C", MatCompositeSetScalings_Composite));
915:   PetscCall(PetscObjectComposeFunction((PetscObject)A, "MatShellSetContext_C", MatShellSetContext_Immutable));
916:   PetscCall(PetscObjectComposeFunction((PetscObject)A, "MatShellSetContextDestroy_C", MatShellSetContextDestroy_Immutable));
917:   PetscCall(PetscObjectComposeFunction((PetscObject)A, "MatShellSetManageScalingShifts_C", MatShellSetManageScalingShifts_Immutable));
918:   PetscCall(PetscObjectChangeTypeName((PetscObject)A, MATCOMPOSITE));
919:   PetscFunctionReturn(PETSC_SUCCESS);
920: }