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*/