Actual source code: ex1.c
1: const char help[] = "TAOTERMCALLBACKS coverage tests";
3: #include <petsctao.h>
5: typedef struct {
6: PetscInt obj_count;
7: PetscInt grad_count;
8: PetscInt obj_and_grad_count;
9: PetscInt hess_count;
10: } AppCtx;
12: static PetscErrorCode objective(Tao tao, Vec x, PetscReal *value, void *ctx)
13: {
14: AppCtx *app = (AppCtx *)ctx;
16: PetscFunctionBeginUser;
17: *value = 0.0;
18: app->obj_count++;
19: PetscFunctionReturn(PETSC_SUCCESS);
20: }
22: static PetscErrorCode gradient(Tao tao, Vec x, Vec g, void *ctx)
23: {
24: AppCtx *app = (AppCtx *)ctx;
26: PetscFunctionBeginUser;
27: PetscCall(VecZeroEntries(g));
28: app->grad_count++;
29: PetscFunctionReturn(PETSC_SUCCESS);
30: }
32: static PetscErrorCode objective_and_gradient(Tao tao, Vec x, PetscReal *value, Vec g, void *ctx)
33: {
34: AppCtx *app = (AppCtx *)ctx;
36: PetscFunctionBeginUser;
37: *value = 0.0;
38: PetscCall(VecZeroEntries(g));
39: app->obj_and_grad_count++;
40: PetscFunctionReturn(PETSC_SUCCESS);
41: }
43: static PetscErrorCode hessian(Tao tao, Vec x, Mat H, Mat Hpre, void *ctx)
44: {
45: AppCtx *app = (AppCtx *)ctx;
47: PetscFunctionBeginUser;
48: if (H) {
49: PetscCall(MatZeroEntries(H));
50: PetscCall(MatAssemblyBegin(H, MAT_FINAL_ASSEMBLY));
51: PetscCall(MatAssemblyEnd(H, MAT_FINAL_ASSEMBLY));
52: }
53: if (Hpre && Hpre != H) {
54: PetscCall(MatZeroEntries(Hpre));
55: PetscCall(MatAssemblyBegin(Hpre, MAT_FINAL_ASSEMBLY));
56: PetscCall(MatAssemblyEnd(Hpre, MAT_FINAL_ASSEMBLY));
57: }
58: app->hess_count++;
59: PetscFunctionReturn(PETSC_SUCCESS);
60: }
62: static PetscErrorCode testCallbacks(PetscBool separate)
63: {
64: Tao tao;
65: TaoTerm term;
66: TaoTermType type;
67: PetscBool same;
68: PetscErrorCode (*_hessian)(Tao, Vec, Mat, Mat, void *);
69: AppCtx app;
70: Vec sol, grad;
71: Mat H, Hpre;
72: PetscInt N = 10;
73: PetscReal value;
75: PetscFunctionBeginUser;
76: app.obj_count = 0;
77: app.grad_count = 0;
78: app.obj_and_grad_count = 0;
79: app.hess_count = 0;
80: PetscCall(VecCreateMPI(PETSC_COMM_WORLD, PETSC_DECIDE, N, &sol));
81: PetscCall(VecDuplicate(sol, &grad));
82: PetscCall(MatCreateAIJ(PETSC_COMM_WORLD, PETSC_DECIDE, PETSC_DECIDE, N, N, 1, NULL, 0, NULL, &H));
83: PetscCall(MatDuplicate(H, MAT_DO_NOT_COPY_VALUES, &Hpre));
84: PetscCall(TaoCreate(PETSC_COMM_WORLD, &tao));
85: PetscCall(TaoSetSolution(tao, sol));
86: PetscCall(TaoGetTerm(tao, NULL, &term, NULL, NULL));
87: PetscCall(TaoTermGetType(term, &type));
88: PetscCall(PetscStrcmp(type, TAOTERMCALLBACKS, &same));
89: PetscCheck(same, PETSC_COMM_WORLD, PETSC_ERR_PLIB, "wrong TaoTermType");
91: if (separate) {
92: PetscCall(TaoSetObjective(tao, objective, (void *)&app));
93: PetscCall(TaoSetGradient(tao, grad, gradient, (void *)&app));
94: } else PetscCall(TaoSetObjectiveAndGradient(tao, grad, objective_and_gradient, (void *)&app));
95: PetscCall(TaoSetHessian(tao, H, Hpre, hessian, (void *)&app));
97: {
98: PetscBool is_defined;
100: PetscCall(TaoTermIsHessianDefined(term, &is_defined));
101: PetscCheck(is_defined == PETSC_TRUE, PETSC_COMM_WORLD, PETSC_ERR_PLIB, "Hessian should be defined after setting it");
102: }
104: if (separate) {
105: PetscErrorCode (*_objective)(Tao, Vec, PetscReal *, void *);
106: PetscErrorCode (*_gradient)(Tao, Vec, Vec, void *);
108: PetscCall(TaoGetObjective(tao, &_objective, NULL));
109: PetscCall(TaoGetGradient(tao, NULL, &_gradient, NULL));
110: PetscCheck(_objective == objective, PETSC_COMM_WORLD, PETSC_ERR_PLIB, "wrong objective callback");
111: PetscCheck(_gradient == gradient, PETSC_COMM_WORLD, PETSC_ERR_PLIB, "wrong gradient callback");
112: } else {
113: PetscErrorCode (*_objective_and_gradient)(Tao, Vec, PetscReal *, Vec, void *);
115: PetscCall(TaoGetObjectiveAndGradient(tao, NULL, &_objective_and_gradient, NULL));
116: PetscCheck(_objective_and_gradient == objective_and_gradient, PETSC_COMM_WORLD, PETSC_ERR_PLIB, "wrong objective and gradient callback");
117: }
118: PetscCall(TaoGetHessian(tao, NULL, NULL, &_hessian, NULL));
119: PetscCheck(_hessian == hessian, PETSC_COMM_WORLD, PETSC_ERR_PLIB, "wrong hessian callback");
121: PetscCall(TaoComputeObjective(tao, sol, &value));
122: (void)value;
123: PetscCall(TaoComputeGradient(tao, sol, grad));
124: PetscCall(TaoComputeObjectiveAndGradient(tao, sol, &value, grad));
125: (void)value;
126: PetscCall(TaoComputeHessian(tao, sol, H, Hpre));
128: if (separate) {
129: PetscCheck(app.obj_count == 2, PETSC_COMM_WORLD, PETSC_ERR_PLIB, "Incorrect number of objective evaluations");
130: PetscCheck(app.grad_count == 2, PETSC_COMM_WORLD, PETSC_ERR_PLIB, "Incorrect number of gradient evaluations");
131: PetscCheck(app.obj_and_grad_count == 0, PETSC_COMM_WORLD, PETSC_ERR_PLIB, "Incorrect number of objective+gradient evaluations");
132: } else {
133: PetscCheck(app.obj_count == 0, PETSC_COMM_WORLD, PETSC_ERR_PLIB, "Incorrect number of objective evaluations");
134: PetscCheck(app.grad_count == 0, PETSC_COMM_WORLD, PETSC_ERR_PLIB, "Incorrect number of gradient evaluations");
135: PetscCheck(app.obj_and_grad_count == 3, PETSC_COMM_WORLD, PETSC_ERR_PLIB, "Incorrect number of objective+gradient evaluations");
136: }
137: PetscCheck(app.hess_count == 1, PETSC_COMM_WORLD, PETSC_ERR_PLIB, "Incorrect number of hessian evaluations");
139: PetscCall(TaoDestroy(&tao));
140: PetscCall(MatDestroy(&Hpre));
141: PetscCall(MatDestroy(&H));
142: PetscCall(VecDestroy(&grad));
143: PetscCall(VecDestroy(&sol));
144: PetscFunctionReturn(PETSC_SUCCESS);
145: }
147: int main(int argc, char **argv)
148: {
149: PetscFunctionBeginUser;
150: PetscCall(PetscInitialize(&argc, &argv, NULL, help));
151: PetscCall(testCallbacks(PETSC_FALSE));
152: PetscCall(testCallbacks(PETSC_TRUE));
153: PetscCall(PetscFinalize());
154: return 0;
155: }
157: /*TEST
159: test:
160: suffix: 0
161: output_file: output/empty.out
163: TEST*/