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