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: }