Actual source code: taotermtest2.c

  1: #include <petsctao.h>

  3: static char help[] = "Using TaoTermShell with mapping matrices that are not diagonal.\n";

  5: typedef struct {
  6:   Vec pdiff_work; /* Work vector for x - params */
  7: } HalfL2Ctx;

  9: typedef struct {
 10:   Mat A;    /* Mapping matrix A */
 11:   Vec p;    /* Target vector p */
 12:   Vec Ax;   /* Work vector for A*x */
 13:   Vec Ax_p; /* Work vector for A*x - p */
 14: } CallbackCtx;

 16: static PetscErrorCode FormFunctionGradient(TaoTerm, Vec, Vec, PetscReal *, Vec);
 17: static PetscErrorCode FormHessian(TaoTerm, Vec, Vec, Mat, Mat);
 18: static PetscErrorCode CtxDestroy(PetscCtxRt ctx);

 20: /* Callback functions for traditional TAO interface */
 21: static PetscErrorCode FormObjectiveGradient_Callback(Tao, Vec, PetscReal *, Vec, void *);
 22: static PetscErrorCode FormHessian_Callback(Tao, Vec, Mat, Mat, void *);

 24: int main(int argc, char **argv)
 25: {
 26:   TaoTerm      objective;
 27:   Tao          tao, tao2;
 28:   PetscMPIInt  size;
 29:   HalfL2Ctx   *ctx;
 30:   MPI_Comm     comm;
 31:   PetscInt     n = 10, m = 10;
 32:   Mat          A;
 33:   Vec          target;
 34:   CallbackCtx *cb_ctx;
 35:   Vec          x_term, x_callback, x2, diff;
 36:   Mat          H2;
 37:   PetscReal    norm_diff, diag_val = 1.1;
 38:   PetscBool    opt, is_diag, is_cdiag, is_aij, is_dense, fd_notpossible;
 39:   const char  *mtype         = MATAIJ;
 40:   char         typeName[256] = "";

 42:   PetscFunctionBeginUser;
 43:   PetscCall(PetscInitialize(&argc, &argv, NULL, help));
 44:   comm = PETSC_COMM_WORLD;
 45:   PetscCallMPI(MPI_Comm_size(comm, &size));
 46:   PetscCheck(size == 1, comm, PETSC_ERR_WRONG_MPI_SIZE, "Incorrect number of processors");

 48:   fd_notpossible = PETSC_FALSE;

 50:   PetscOptionsBegin(comm, "", help, "none");
 51:   PetscCall(PetscOptionsBool("-fd_notpossible", "Set TaoTermShell ComputeHessianFDPossible as false", "", fd_notpossible, &fd_notpossible, NULL));
 52:   PetscCall(PetscOptionsInt("-n", "Problem size", "", n, &n, NULL));
 53:   PetscCall(PetscOptionsInt("-m", "Mapping matrix row size", "", m, &m, NULL));
 54:   PetscCall(PetscOptionsReal("-diag_val", "Value of constant diagonal matrix", NULL, diag_val, &diag_val, NULL));
 55:   PetscCall(PetscOptionsFList("-mapping_mtype", "Mapping matrix type", "", MatList, mtype, typeName, sizeof(typeName), &opt));
 56:   PetscOptionsEnd();

 58:   PetscCall(PetscNew(&ctx));

 60:   /* Initialize typeName to default if option was not set */
 61:   if (!opt) PetscCall(PetscStrcpy(typeName, mtype));

 63:   PetscCall(PetscStrcmp(typeName, MATDIAGONAL, &is_diag));
 64:   PetscCall(PetscStrcmp(typeName, MATCONSTANTDIAGONAL, &is_cdiag));
 65:   PetscCall(PetscStrcmp(typeName, MATAIJ, &is_aij));
 66:   PetscCall(PetscStrcmp(typeName, MATDENSE, &is_dense));
 67:   /* Create mapping matrix A: m x n (maps from solution space to term space) */
 68:   if (is_diag) {
 69:     /* Create a diagonal matrix */
 70:     Vec      diag_vec;
 71:     PetscInt diag_size;

 73:     PetscCheck(m == n, comm, PETSC_ERR_ARG_INCOMP, "For diagonal matrix, m and n must be equal (got m=%" PetscInt_FMT ", n=%" PetscInt_FMT ")", m, n);
 74:     diag_size = m;
 75:     PetscCall(VecCreate(comm, &diag_vec));
 76:     PetscCall(VecSetSizes(diag_vec, PETSC_DECIDE, diag_size));
 77:     PetscCall(VecSetFromOptions(diag_vec));
 78:     PetscCall(VecSetRandom(diag_vec, NULL));
 79:     PetscCall(MatCreateDiagonal(diag_vec, &A));
 80:     PetscCall(VecDestroy(&diag_vec));
 81:   } else if (is_cdiag) {
 82:     /* Create a constant diagonal matrix */
 83:     PetscCheck(m == n, comm, PETSC_ERR_ARG_INCOMP, "For constant diagonal matrix, m and n must be equal (got m=%" PetscInt_FMT ", n=%" PetscInt_FMT ")", m, n);
 84:     PetscCall(MatCreateConstantDiagonal(comm, PETSC_DECIDE, PETSC_DECIDE, m, n, diag_val, &A));
 85:   } else if (is_dense) {
 86:     /* Create a dense matrix */
 87:     PetscCall(MatCreateDense(comm, PETSC_DECIDE, PETSC_DECIDE, m, n, NULL, &A));
 88:     PetscCall(MatSetFromOptions(A));
 89:     PetscCall(MatSetRandom(A, NULL));
 90:     PetscCall(MatAssemblyBegin(A, MAT_FINAL_ASSEMBLY));
 91:     PetscCall(MatAssemblyEnd(A, MAT_FINAL_ASSEMBLY));
 92:   } else {
 93:     /* Create an AIJ matrix (default) */
 94:     PetscCall(MatCreateSeqAIJ(comm, m, n, PETSC_DEFAULT, NULL, &A));
 95:     PetscCall(MatSetFromOptions(A));
 96:     PetscCall(MatSetRandom(A, NULL));
 97:     PetscCall(MatAssemblyBegin(A, MAT_FINAL_ASSEMBLY));
 98:     PetscCall(MatAssemblyEnd(A, MAT_FINAL_ASSEMBLY));
 99:   }

101:   /* Create shell term that computes f(x) = 0.5 ||x||_2^2 */
102:   PetscCall(TaoTermCreateShell(comm, ctx, CtxDestroy, &objective));

104:   /* Set solution and parameter sizes to match the mapped space (m) */
105:   PetscCall(TaoTermSetSolutionSizes(objective, PETSC_DECIDE, m, 1));
106:   PetscCall(TaoTermSetParametersSizes(objective, PETSC_DECIDE, m, 1));

108:   PetscCall(TaoTermShellSetObjectiveAndGradient(objective, FormFunctionGradient));
109:   PetscCall(TaoTermShellSetCreateHessianMatrices(objective, TaoTermCreateHessianMatricesDefault));
110:   PetscCall(TaoTermSetCreateHessianMode(objective, PETSC_TRUE /* H == Hpre */, MATAIJ, NULL));
111:   PetscCall(TaoTermShellSetHessian(objective, FormHessian));
112:   PetscCall(TaoTermSetFromOptions(objective));
113:   if (fd_notpossible) PetscCall(TaoTermShellSetIsComputeHessianFDPossible(objective, PETSC_BOOL3_FALSE));

115:   PetscCall(TaoTermSetUp(objective));

117:   /* Create target vector for least squares problem (parameters) */
118:   PetscCall(TaoTermCreateParametersVec(objective, &target));
119:   PetscCall(VecSetRandom(target, NULL));

121:   PetscCall(TaoCreate(comm, &tao));
122:   PetscCall(PetscObjectSetOptionsPrefix((PetscObject)tao, "shell_"));
123:   PetscCall(TaoSetType(tao, TAOLMVM));

125:   /* Add term with mapping matrix A: f(Ax; p) = 0.5 ||Ax - p||_2^2 */
126:   PetscCall(TaoAddTerm(tao, NULL, 1.0, objective, target, A));

128:   PetscCall(TaoSetFromOptions(tao));
129:   PetscCall(TaoSolve(tao));

131:   /* Allocate callback context */
132:   PetscCall(PetscNew(&cb_ctx));
133:   cb_ctx->A = A;
134:   cb_ctx->p = target;

136:   /* Create work vectors */
137:   PetscCall(MatCreateVecs(A, NULL, &cb_ctx->Ax));
138:   PetscCall(VecDuplicate(target, &cb_ctx->Ax_p));

140:   PetscCall(MatCreateVecs(A, &x2, NULL));

142:   /* Create Hessian matrix A^T * A */
143:   if (is_diag) {
144:     Vec A_diag, H2_diag;

146:     PetscCall(MatCreateVecs(A, &A_diag, NULL));
147:     PetscCall(MatGetDiagonal(A, A_diag));
148:     PetscCall(VecDuplicate(A_diag, &H2_diag));
149:     PetscCall(VecPointwiseMult(H2_diag, A_diag, A_diag));
150:     PetscCall(MatCreateDiagonal(H2_diag, &H2));
151:     PetscCall(VecDestroy(&A_diag));
152:     PetscCall(VecDestroy(&H2_diag));
153:   } else if (is_cdiag) {
154:     PetscCall(MatCreateConstantDiagonal(comm, PETSC_DECIDE, PETSC_DECIDE, m, n, diag_val * diag_val, &H2));
155:   } else {
156:     Mat       Htest, Hpretest;
157:     PetscBool is_h_dense;

159:     PetscCall(MatTransposeMatMult(A, A, MAT_INITIAL_MATRIX, PETSC_DETERMINE, &H2));
160:     PetscCall(MatAssemblyBegin(H2, MAT_FINAL_ASSEMBLY));
161:     PetscCall(MatAssemblyEnd(H2, MAT_FINAL_ASSEMBLY));

163:     PetscCall(TaoGetHessianMatrices(tao, &Htest, &Hpretest));
164:     PetscCall(PetscObjectBaseTypeCompare((PetscObject)Htest, MATSEQDENSE, &is_h_dense));
165:     if (is_h_dense) PetscCall(MatConvert(H2, MATDENSE, MAT_INPLACE_MATRIX, &H2));
166:   }
167:   /* Create second TAO solver */
168:   PetscCall(TaoCreate(comm, &tao2));
169:   PetscCall(PetscObjectSetOptionsPrefix((PetscObject)tao2, "regular_"));
170:   PetscCall(TaoSetType(tao2, TAOLMVM));
171:   PetscCall(TaoSetSolution(tao2, x2));
172:   PetscCall(TaoSetObjectiveAndGradient(tao2, NULL, FormObjectiveGradient_Callback, cb_ctx));
173:   PetscCall(TaoSetHessian(tao2, H2, H2, FormHessian_Callback, cb_ctx));
174:   PetscCall(TaoSetFromOptions(tao2));
175:   PetscCall(TaoSolve(tao2));

177:   /* Compare solutions */
178:   PetscCall(TaoGetSolution(tao, &x_term));
179:   PetscCall(TaoGetSolution(tao2, &x_callback));
180:   PetscCall(VecDuplicate(x_term, &diff));
181:   PetscCall(VecCopy(x_term, diff));
182:   PetscCall(VecAXPY(diff, -1.0, x_callback));
183:   PetscCall(VecNorm(diff, NORM_2, &norm_diff));
184:   if (norm_diff <= 1.e-12) PetscCall(PetscPrintf(comm, "Relative difference < 1e-12\n"));
185:   else PetscCall(PetscPrintf(comm, "Relative difference > 1e-12: %6.10e\n", (double)norm_diff));
186:   PetscCall(VecDestroy(&x2));
187:   PetscCall(VecDestroy(&diff));
188:   PetscCall(VecDestroy(&cb_ctx->Ax));
189:   PetscCall(VecDestroy(&cb_ctx->Ax_p));
190:   PetscCall(PetscFree(cb_ctx));
191:   PetscCall(VecDestroy(&target));
192:   PetscCall(MatDestroy(&A));
193:   PetscCall(MatDestroy(&H2));
194:   PetscCall(TaoDestroy(&tao2));
195:   PetscCall(TaoDestroy(&tao));
196:   PetscCall(TaoTermDestroy(&objective));
197:   PetscCall(PetscFinalize());
198:   return 0;
199: }

201: /*
202:   FormFunctionGradient - Evaluates the function, f(X), and gradient, G(X).

204:   Input Parameters:
205: + term      - the `TaoTerm` for the objective function
206: . x         - input vector
207: - params    - optional vector of parameters

209:   Output Parameters:
210: + f - function value
211: - G - vector containing the newly evaluated gradient

213:   Note:
214:   Computes f = 0.5 * ||x - params||_2^2 and g = x - params, matching TAOTERMHALFL2SQUARED.
215: */
216: static PetscErrorCode FormFunctionGradient(TaoTerm term, Vec x, Vec params, PetscReal *f, Vec G)
217: {
218:   HalfL2Ctx  *ctx;
219:   PetscScalar v;

221:   PetscFunctionBeginUser;
222:   PetscCall(TaoTermShellGetContext(term, &ctx));
223:   if (params) {
224:     PetscCall(VecWAXPY(G, -1.0, params, x));
225:     PetscCall(VecDot(G, G, &v));
226:   } else {
227:     PetscCall(VecCopy(x, G));
228:     PetscCall(VecDot(G, G, &v));
229:   }
230:   *f = 0.5 * PetscRealPart(v);
231:   PetscFunctionReturn(PETSC_SUCCESS);
232: }

234: /*
235:   FormHessian - Evaluates Hessian matrix.

237:   Input Parameters:
238: + term      - the `TaoTerm` for the objective function
239: . x         - input vector
240: . params    - optional vector of parameters
241: - Hpre      - optional matrix for building the preconditioner

243:   Output Parameters:
244: + H    - Hessian matrix
245: - Hpre - matrix for building the preconditioning

247:   Note:
248:   Computes H = I (identity matrix), matching TAOTERMHALFL2SQUARED.
249: */
250: static PetscErrorCode FormHessian(TaoTerm term, Vec x, Vec params, Mat H, Mat Hpre)
251: {
252:   PetscFunctionBeginUser;
253:   if (H) {
254:     PetscCall(MatZeroEntries(H));
255:     PetscCall(MatAssemblyBegin(H, MAT_FINAL_ASSEMBLY));
256:     PetscCall(MatAssemblyEnd(H, MAT_FINAL_ASSEMBLY));
257:     PetscCall(MatShift(H, 1.0));
258:   }
259:   if (Hpre && Hpre != H) {
260:     PetscCall(MatZeroEntries(Hpre));
261:     PetscCall(MatAssemblyBegin(Hpre, MAT_FINAL_ASSEMBLY));
262:     PetscCall(MatAssemblyEnd(Hpre, MAT_FINAL_ASSEMBLY));
263:     PetscCall(MatShift(Hpre, 1.0));
264:   }
265:   PetscFunctionReturn(PETSC_SUCCESS);
266: }

268: static PetscErrorCode CtxDestroy(PetscCtxRt ctx_ptr)
269: {
270:   HalfL2Ctx *ctx = *(HalfL2Ctx **)ctx_ptr;

272:   PetscFunctionBeginUser;
273:   if (ctx) {
274:     PetscCall(VecDestroy(&ctx->pdiff_work));
275:     PetscCall(PetscFree(ctx));
276:     *(void **)ctx_ptr = NULL;
277:   }
278:   PetscFunctionReturn(PETSC_SUCCESS);
279: }

281: /*
282:   FormObjectiveGradient_Callback - Evaluates the objective and gradient for traditional TAO callback interface.

284:   Input Parameters:
285: + tao  - the Tao solver context
286: . x    - input vector (size n)
287: - ctx  - application context containing A and p

289:   Output Parameters:
290: + f - function value: 0.5 * ||Ax - p||_2^2
291: - g - gradient vector: A^T (Ax - p)

293:   Note:
294:   Computes f = 0.5 * ||Ax - p||_2^2 and g = A^T (Ax - p)
295: */
296: static PetscErrorCode FormObjectiveGradient_Callback(Tao tao, Vec x, PetscReal *f, Vec g, void *ctx)
297: {
298:   CallbackCtx *cb_ctx = (CallbackCtx *)ctx;
299:   PetscScalar  v;

301:   PetscFunctionBeginUser;
302:   /* Compute Ax */
303:   PetscCall(MatMult(cb_ctx->A, x, cb_ctx->Ax));
304:   /* Compute Ax - p */
305:   PetscCall(VecCopy(cb_ctx->Ax, cb_ctx->Ax_p));
306:   PetscCall(VecAXPY(cb_ctx->Ax_p, -1.0, cb_ctx->p));
307:   /* Compute objective: 0.5 * ||Ax - p||_2^2 */
308:   PetscCall(VecDot(cb_ctx->Ax_p, cb_ctx->Ax_p, &v));
309:   *f = 0.5 * PetscRealPart(v);
310:   /* Compute gradient: A^T (Ax - p) */
311:   PetscCall(MatMultTranspose(cb_ctx->A, cb_ctx->Ax_p, g));
312:   PetscFunctionReturn(PETSC_SUCCESS);
313: }

315: /*
316:   FormHessian_Callback - Evaluates the Hessian matrix for traditional TAO callback interface.

318:   Input Parameters:
319: + tao  - the Tao solver context
320: . x    - input vector
321: . H    - Hessian matrix (should be pre-allocated as A^T * A)
322: . Hpre - preconditioner matrix
323: - ctx  - application context containing A and p

325:   Output Parameters:
326: + H    - Hessian matrix (A^T * A)
327: - Hpre - Preconditioning matrix

329:   Note:
330:   The Hessian for 0.5 * ||Ax - p||_2^2 is constant: H = A^T * A
331: */
332: static PetscErrorCode FormHessian_Callback(Tao tao, Vec x, Mat H, Mat Hpre, void *ctx)
333: {
334:   PetscFunctionBeginUser;
335:   /* Hessian is constant: A^T * A, which should already be set in H */
336:   if (Hpre && Hpre != H) PetscCall(MatCopy(H, Hpre, SAME_NONZERO_PATTERN));
337:   PetscFunctionReturn(PETSC_SUCCESS);
338: }

340: /* Note: For dense variations, relative error may be greater than 1.e-12, *
341:  * but that is okay, as it is a result of KSP, and PC using AIJ matrices  *
342:  * instead of dense.                                                      */

344: /*TEST

346:    build:
347:      requires: !complex !single !quad !defined(PETSC_USE_64BIT_INDICES) !__float128

349:    test:
350:      suffix: diag_diag
351:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
352:      args: -tao_term_hessian_mat_type diagonal -mapping_mtype diagonal

354:    test:
355:      suffix: diag_cdiag
356:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
357:      args: -tao_term_hessian_mat_type diagonal -mapping_mtype constantdiagonal

359:    test:
360:      suffix: diag_dense
361:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
362:      args: -tao_term_hessian_mat_type diagonal -mapping_mtype dense

364:    test:
365:      suffix: diag_dense_nsq
366:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
367:      args: -tao_term_hessian_mat_type diagonal -mapping_mtype dense -m 15

369:    test:
370:      suffix: diag_aij
371:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
372:      args: -tao_term_hessian_mat_type diagonal -mapping_mtype aij

374:    test:
375:      suffix: cdiag_diag
376:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
377:      args: -tao_term_hessian_mat_type constantdiagonal -mapping_mtype diagonal

379:    test:
380:      suffix: cdiag_cdiag
381:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
382:      args: -tao_term_hessian_mat_type constantdiagonal -mapping_mtype constantdiagonal

384:    test:
385:      suffix: cdiag_dense
386:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
387:      args: -tao_term_hessian_mat_type constantdiagonal -mapping_mtype dense

389:    test:
390:      suffix: cdiag_dense_nsq
391:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
392:      args: -tao_term_hessian_mat_type constantdiagonal -mapping_mtype dense -m 15

394:    test:
395:      suffix: cdiag_aij
396:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
397:      args: -tao_term_hessian_mat_type constantdiagonal -mapping_mtype aij

399:    test:
400:      suffix: dense_diag
401:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
402:      args: -tao_term_hessian_mat_type dense -mapping_mtype diagonal

404:    test:
405:      suffix: dense_cdiag
406:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
407:      args: -tao_term_hessian_mat_type dense -mapping_mtype constantdiagonal

409:    test:
410:      suffix: dense_dense
411:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
412:      args: -tao_term_hessian_mat_type dense -mapping_mtype dense -fd_notpossible {{0 1}}

414:    test:
415:      suffix: dense_dense_nsq
416:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
417:      args: -tao_term_hessian_mat_type dense -mapping_mtype dense -m 15

419:    test:
420:      suffix: dense_aij
421:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
422:      args: -tao_term_hessian_mat_type dense -mapping_mtype aij

424:    test:
425:      suffix: aij_diag
426:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
427:      args: -tao_term_hessian_mat_type aij -mapping_mtype diagonal

429:    test:
430:      suffix: aij_cdiag
431:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
432:      args: -tao_term_hessian_mat_type aij -mapping_mtype constantdiagonal

434:    test:
435:      suffix: aij_dense
436:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
437:      args: -tao_term_hessian_mat_type aij -mapping_mtype dense

439:    test:
440:      suffix: aij_dense_nsq
441:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
442:      args: -tao_term_hessian_mat_type aij -mapping_mtype dense -m 15

444:    test:
445:      suffix: aij_aij
446:      args: -shell_tao_type nls -shell_tao_view ::ascii_info_detail -regular_tao_type nls -regular_tao_view ::ascii_info_detail
447:      args: -tao_term_hessian_mat_type aij -mapping_mtype aij

449: TEST*/