Actual source code: letkf_obs_scatter.c

  1: #include <petsc.h>
  2: #include <petsc/private/hashmapi.h>
  3: #include <../src/ml/da/impls/ensemble/letkf/letkf.h>

  5: /*
  6:   PetscDALETKFDestroyObsScatter - Release the IS, hash, scatter context, and local work vectors
  7:   built by PetscDALETKFSetupObsScatter(). Idempotent; safe to call when nothing has been set up.
  8: */
  9: PETSC_INTERN PetscErrorCode PetscDALETKFDestroyObsScatter(PetscDA_LETKF *impl)
 10: {
 11:   PetscFunctionBegin;
 12:   PetscCall(ISDestroy(&impl->obs_is_local));
 13:   PetscCall(VecScatterDestroy(&impl->obs_scat));
 14:   PetscCall(VecDestroy(&impl->obs_work));
 15:   PetscCall(VecDestroy(&impl->y_mean_work));
 16:   PetscCall(VecDestroy(&impl->r_inv_sqrt_work));
 17:   PetscCall(MatDestroy(&impl->Z_work));
 18:   PetscCall(PetscHMapIDestroy(&impl->obs_g2l));
 19:   PetscFunctionReturn(PETSC_SUCCESS);
 20: }

 22: /*
 23:   PetscDALETKFSetupObsScatter - Build the IS, global-to-local hash, scatter context, and local
 24:   work vectors needed by the per-vertex CPU and Kokkos analysis paths.

 26:   For each row of impl->Q owned by this rank, collect the unique global observation indices
 27:   referenced by that row's column indices. Sort them, build an IS, and create a VecScatter from
 28:   any global observation-space vector (template taken from H) into a sequential work vector of
 29:   matching length. Also build a hash mapping global obs index -> local position so the per-vertex
 30:   extractor can translate column indices on the fly.

 32:   Backend-agnostic: contains no Kokkos calls. Both the CPU analysis path and the device-CSR setup
 33:   in PetscDALETKFSetupLocalization_Kokkos() consume what this populates.
 34: */
 35: PETSC_INTERN PetscErrorCode PetscDALETKFSetupObsScatter(PetscDA_LETKF *impl, Mat H)
 36: {
 37:   PetscInt      rstart, rend, nrows, n_obs_global_Q, n_obs_global_H;
 38:   PetscInt      n_obs_local_total = 0, off = 0;
 39:   PetscInt     *obs_indices;
 40:   PetscHashIter iter;
 41:   PetscBool     missing;
 42:   Vec           gvec;
 43:   IS            is_to;

 45:   PetscFunctionBegin;
 46:   PetscCheck(impl->Q, PetscObjectComm((PetscObject)H), PETSC_ERR_ARG_WRONGSTATE, "impl->Q must be installed before PetscDALETKFSetupObsScatter()");

 48:   /* Q's columns are global obs indices; H supplies the source layout for the scatter. They
 49:      must agree on the global obs-space size, else the scatter would dereference out-of-range
 50:      global indices. Catches structurally-mismatched H reaching analysis without invalidating Q. */
 51:   PetscCall(MatGetSize(impl->Q, NULL, &n_obs_global_Q));
 52:   PetscCall(MatGetSize(H, &n_obs_global_H, NULL));
 53:   PetscCheck(n_obs_global_Q == n_obs_global_H, PetscObjectComm((PetscObject)H), PETSC_ERR_ARG_INCOMP, "Q columns (%" PetscInt_FMT ") and H rows (%" PetscInt_FMT ") must agree on global obs-space size; re-supply coordinates if H changed", n_obs_global_Q, n_obs_global_H);

 55:   PetscCall(MatGetOwnershipRange(impl->Q, &rstart, &rend));
 56:   nrows = rend - rstart;

 58:   PetscCall(PetscHMapICreate(&impl->obs_g2l));

 60:   /* First pass: insert every column index into the hashmap to deduplicate. Sizing the workspace
 61:      by the unique count (rather than the sum-of-row-nnz upper bound used previously) keeps peak
 62:      memory O(unique obs) instead of O(sum nnz). */
 63:   for (PetscInt i = 0; i < nrows; i++) {
 64:     const PetscInt *cols;
 65:     PetscInt        nnz;

 67:     PetscCall(MatGetRow(impl->Q, rstart + i, &nnz, &cols, NULL));
 68:     for (PetscInt k = 0; k < nnz; k++) PetscCall(PetscHMapIPut(impl->obs_g2l, cols[k], &iter, &missing));
 69:     PetscCall(MatRestoreRow(impl->Q, rstart + i, &nnz, &cols, NULL));
 70:   }
 71:   PetscCall(PetscHMapIGetSize(impl->obs_g2l, &n_obs_local_total));

 73:   PetscCall(PetscMalloc1(n_obs_local_total, &obs_indices));
 74:   PetscCall(PetscHMapIGetKeys(impl->obs_g2l, &off, obs_indices));
 75:   PetscCall(PetscSortInt(n_obs_local_total, obs_indices));

 77:   PetscCall(ISCreateGeneral(PETSC_COMM_SELF, n_obs_local_total, obs_indices, PETSC_COPY_VALUES, &impl->obs_is_local));

 79:   /* Repopulate obs_g2l with sorted-position values: each global index maps to its slot in obs_work
 80:      after the scatter. */
 81:   PetscCall(PetscHMapIClear(impl->obs_g2l));
 82:   for (PetscInt i = 0; i < n_obs_local_total; i++) PetscCall(PetscHMapISet(impl->obs_g2l, obs_indices[i], i));

 84:   PetscCall(PetscFree(obs_indices));

 86:   PetscCall(VecCreateSeq(PETSC_COMM_SELF, n_obs_local_total, &impl->obs_work));
 87:   PetscCall(VecDuplicate(impl->obs_work, &impl->y_mean_work));
 88:   PetscCall(VecDuplicate(impl->obs_work, &impl->r_inv_sqrt_work));

 90:   PetscCall(MatCreateVecs(H, NULL, &gvec));
 91:   PetscCall(ISCreateStride(PETSC_COMM_SELF, n_obs_local_total, 0, 1, &is_to));
 92:   PetscCall(VecScatterCreate(gvec, impl->obs_is_local, impl->obs_work, is_to, &impl->obs_scat));
 93:   PetscCall(VecDestroy(&gvec));
 94:   PetscCall(ISDestroy(&is_to));
 95:   PetscFunctionReturn(PETSC_SUCCESS);
 96: }