Actual source code: dalocalizationletkf.c
1: #include <petsc.h>
2: #include <petscmat.h>
3: #include <../src/ml/da/impls/ensemble/letkf/letkf.h>
4: #include <../src/ml/da/impls/ensemble/letkf/letkf_kernels.h>
6: /* Bbox layout is `[min_0..min_{dim-1}, max_0..max_{dim-1}]`, i.e. all mins then all maxes. */
7: static inline PetscBool PetscDALETKFCoordInBbox(PetscInt dim, const PetscReal *coord, const PetscReal *bbox)
8: {
9: for (PetscInt d = 0; d < dim; ++d)
10: if (coord[d] < bbox[d] || coord[d] > bbox[dim + d]) return PETSC_FALSE;
11: return PETSC_TRUE;
12: }
14: /*
15: PetscDALETKFGatherObsBbox - Bounding-box-pruned redistribution of observation coordinates.
17: Each rank obtains the global indices and coordinates of just those observations whose location can
18: fall within `cutoff` of any vertex it owns. Replaces an earlier all-gather-to-all of the full obs
19: coordinate set, which scaled as O(n_obs_global) per rank in both memory and bandwidth.
21: For each non-periodic dimension d, the local-vertex bbox is `[vmin_d - cutoff, vmax_d + cutoff]`.
22: Bboxes are exchanged with `MPI_Allgather`; each rank then scans its locally-owned obs (in `H`'s row
23: distribution) and forwards every obs into each rank whose padded bbox contains it. Periodic
24: dimensions (`bd[d] > 0`) are passed unfiltered because wrap-around defeats a scalar bbox test;
25: the exact distance gate inside the Q-construction kernels handles those dims.
27: Output buffers `*obs_idx_out` (global obs indices, length `*n_obs_filt`) and `*obs_coords_out`
28: (flattened `[*n_obs_filt][dim]` row-major) are allocated with `PetscMalloc1()` and must be freed
29: by the caller. The order of records is whatever Alltoallv produces - column ordering inside Q is
30: re-sorted by `MatSetValues()`, so callers may use it directly.
32: The per-row distance check inside the AIJ and Kokkos Q-construction kernels still gates on weight
33: > 0, so the resulting Q is identical to the all-gather variant - this routine only trims the
34: candidate set passed in.
35: */
36: PETSC_INTERN PetscErrorCode PetscDALETKFGatherObsBbox(PetscInt dim, Vec xyz[], PetscReal bd[], PetscReal cutoff, Mat H, Vec obs_vecs[], PetscInt *n_obs_filt, PetscInt **obs_idx_out, PetscReal **obs_coords_out)
37: {
38: MPI_Comm comm;
39: PetscMPIInt size, two_dim_mpi;
40: PetscInt n_vert_local, obs_rstart, obs_rend, n_obs_local;
41: PetscInt n_periodic = 0;
42: PetscInt total_send = 0, total_recv = 0;
43: PetscReal local_bbox[2 * 3] = {0.0, 0.0, 0.0, 0.0, 0.0, 0.0}; /* layout: [vmin_0..vmin_{d-1}, vmax_0..vmax_{d-1}] */
44: PetscReal *all_bboxes = NULL;
45: PetscInt *send_counts = NULL, *send_displs = NULL, *recv_counts = NULL, *recv_displs = NULL;
46: PetscInt *send_idx = NULL, *recv_idx = NULL, *pos = NULL;
47: PetscReal *send_crd = NULL, *recv_crd = NULL;
48: PetscMPIInt *send_counts_mpi = NULL, *send_displs_mpi = NULL, *recv_counts_mpi = NULL, *recv_displs_mpi = NULL;
49: PetscMPIInt *send_counts_crd = NULL, *send_displs_crd = NULL, *recv_counts_crd = NULL, *recv_displs_crd = NULL;
50: const PetscScalar *obs_arr[3] = {NULL, NULL, NULL};
52: PetscFunctionBegin;
53: PetscCall(PetscObjectGetComm((PetscObject)H, &comm));
54: PetscCheck(dim >= 1 && dim <= 3, comm, PETSC_ERR_ARG_OUTOFRANGE, "Spatial dimension must be in [1, 3]; got %" PetscInt_FMT " (local_bbox is sized for at most 3 dims)", dim);
55: PetscCallMPI(MPI_Comm_size(comm, &size));
56: PetscCall(PetscMPIIntCast(2 * dim, &two_dim_mpi));
57: PetscCall(VecGetLocalSize(xyz[0], &n_vert_local));
58: PetscCall(MatGetOwnershipRange(H, &obs_rstart, &obs_rend));
59: n_obs_local = obs_rend - obs_rstart;
61: /* Single-rank fast path: no exchange needed, and MPI-uni does not implement Alltoallv. */
62: if (size == 1) {
63: PetscInt *idx_out;
64: PetscReal *crd_out;
66: PetscCall(PetscMalloc1(n_obs_local, &idx_out));
67: PetscCall(PetscMalloc1((size_t)n_obs_local * dim, &crd_out));
68: for (PetscInt d = 0; d < dim; ++d) PetscCall(VecGetArrayRead(obs_vecs[d], &obs_arr[d]));
69: for (PetscInt k = 0; k < n_obs_local; ++k) {
70: idx_out[k] = obs_rstart + k;
71: for (PetscInt d = 0; d < dim; ++d) crd_out[(size_t)k * dim + d] = PetscRealPart(obs_arr[d][k]);
72: }
73: for (PetscInt d = 0; d < dim; ++d) PetscCall(VecRestoreArrayRead(obs_vecs[d], &obs_arr[d]));
74: *n_obs_filt = n_obs_local;
75: *obs_idx_out = idx_out;
76: *obs_coords_out = crd_out;
77: PetscFunctionReturn(PETSC_SUCCESS);
78: }
80: /* Step 1: per-dim local-vertex bbox padded by cutoff. Periodic dims pass through (-/+ MAX_REAL).
81: A rank with no vertices uses (+MAX, -MAX) so its bbox excludes every obs. When every dim is
82: periodic the prune is a no-op and the gather degenerates to all-to-all; emit a PetscInfo
83: under `-info` so this is visible without surprising serial users. Fires on every Q rebuild
84: (Q is built lazily and re-built when the setters invalidate it). */
85: for (PetscInt d = 0; d < dim; ++d)
86: if (bd[d] > 0.0) n_periodic++;
87: /* size == 1 already returned above, so the prune-is-a-no-op notice always applies to a real multi-rank gather. */
88: if (n_periodic == dim) PetscCall(PetscInfo((PetscObject)H, "All %" PetscInt_FMT " dim(s) periodic; bbox prune passes through, obs gather degenerates to all-to-all\n", dim));
89: for (PetscInt d = 0; d < dim; ++d) {
90: if (bd[d] > 0.0) {
91: local_bbox[d] = -PETSC_MAX_REAL;
92: local_bbox[dim + d] = PETSC_MAX_REAL;
93: } else if (n_vert_local == 0) {
94: local_bbox[d] = PETSC_MAX_REAL;
95: local_bbox[dim + d] = -PETSC_MAX_REAL;
96: } else {
97: const PetscScalar *arr;
98: PetscReal vmin = PETSC_MAX_REAL, vmax = -PETSC_MAX_REAL;
100: PetscCall(VecGetArrayRead(xyz[d], &arr));
101: for (PetscInt i = 0; i < n_vert_local; ++i) {
102: PetscReal x = PetscRealPart(arr[i]);
103: if (x < vmin) vmin = x;
104: if (x > vmax) vmax = x;
105: }
106: PetscCall(VecRestoreArrayRead(xyz[d], &arr));
107: local_bbox[d] = vmin - cutoff;
108: local_bbox[dim + d] = vmax + cutoff;
109: }
110: }
112: /* Step 2: Allgather bboxes. Layout per rank: [vmin_0..vmin_{d-1}, vmax_0..vmax_{d-1}]. */
113: PetscCall(PetscMalloc1((size_t)size * 2 * dim, &all_bboxes));
114: PetscCallMPI(MPI_Allgather(local_bbox, two_dim_mpi, MPIU_REAL, all_bboxes, two_dim_mpi, MPIU_REAL, comm));
116: /* Step 3: Count records per dest rank. The Get/Restore pair is opened only over the count loop so
117: a malloc failure later cannot leave obs_vecs locked. */
118: PetscCall(PetscMalloc2(size, &send_counts, size + 1, &send_displs));
119: PetscCall(PetscArrayzero(send_counts, size));
120: for (PetscInt d = 0; d < dim; ++d) PetscCall(VecGetArrayRead(obs_vecs[d], &obs_arr[d]));
121: for (PetscInt k = 0; k < n_obs_local; ++k) {
122: PetscReal coord_k[3] = {0.0, 0.0, 0.0};
124: for (PetscInt d = 0; d < dim; ++d) coord_k[d] = PetscRealPart(obs_arr[d][k]);
125: for (PetscMPIInt r = 0; r < size; ++r) {
126: if (PetscDALETKFCoordInBbox(dim, coord_k, &all_bboxes[(size_t)r * 2 * dim])) send_counts[r]++;
127: }
128: }
129: for (PetscInt d = 0; d < dim; ++d) PetscCall(VecRestoreArrayRead(obs_vecs[d], &obs_arr[d]));
131: send_displs[0] = 0;
132: for (PetscMPIInt r = 0; r < size; ++r) send_displs[r + 1] = send_displs[r] + send_counts[r];
133: total_send = send_displs[size];
135: /* Step 4: Pack send buffers (re-acquire obs_vecs only for the duration of the pack). */
136: PetscCall(PetscMalloc1(total_send, &send_idx));
137: PetscCall(PetscMalloc1((size_t)total_send * dim, &send_crd));
138: PetscCall(PetscMalloc1(size, &pos));
139: for (PetscMPIInt r = 0; r < size; ++r) pos[r] = send_displs[r];
140: for (PetscInt d = 0; d < dim; ++d) PetscCall(VecGetArrayRead(obs_vecs[d], &obs_arr[d]));
141: for (PetscInt k = 0; k < n_obs_local; ++k) {
142: PetscReal coord_k[3] = {0.0, 0.0, 0.0};
144: for (PetscInt d = 0; d < dim; ++d) coord_k[d] = PetscRealPart(obs_arr[d][k]);
145: for (PetscMPIInt r = 0; r < size; ++r) {
146: if (PetscDALETKFCoordInBbox(dim, coord_k, &all_bboxes[(size_t)r * 2 * dim])) {
147: PetscInt p = pos[r]++;
148: send_idx[p] = obs_rstart + k;
149: for (PetscInt d = 0; d < dim; ++d) send_crd[(size_t)p * dim + d] = coord_k[d];
150: }
151: }
152: }
153: for (PetscInt d = 0; d < dim; ++d) PetscCall(VecRestoreArrayRead(obs_vecs[d], &obs_arr[d]));
154: PetscCall(PetscFree(pos));
156: /* Step 5: Exchange counts, then index and coord payloads. */
157: PetscCall(PetscMalloc2(size, &recv_counts, size + 1, &recv_displs));
158: PetscCall(PetscMalloc4(size, &send_counts_mpi, size, &send_displs_mpi, size, &recv_counts_mpi, size, &recv_displs_mpi));
159: for (PetscMPIInt r = 0; r < size; ++r) {
160: PetscCall(PetscMPIIntCast(send_counts[r], &send_counts_mpi[r]));
161: PetscCall(PetscMPIIntCast(send_displs[r], &send_displs_mpi[r]));
162: }
163: PetscCallMPI(MPI_Alltoall(send_counts_mpi, 1, MPI_INT, recv_counts_mpi, 1, MPI_INT, comm));
164: recv_displs_mpi[0] = 0;
165: for (PetscMPIInt r = 1; r < size; ++r) PetscCall(PetscMPIIntCast((PetscInt64)recv_displs_mpi[r - 1] + (PetscInt64)recv_counts_mpi[r - 1], &recv_displs_mpi[r]));
166: recv_displs[0] = 0;
167: for (PetscMPIInt r = 0; r < size; ++r) {
168: recv_counts[r] = recv_counts_mpi[r];
169: recv_displs[r + 1] = recv_displs[r] + recv_counts[r];
170: }
171: total_recv = recv_displs[size];
173: PetscCall(PetscMalloc1(total_recv, &recv_idx));
174: PetscCall(PetscMalloc1((size_t)total_recv * dim, &recv_crd));
176: PetscCallMPI(MPI_Alltoallv(send_idx, send_counts_mpi, send_displs_mpi, MPIU_INT, recv_idx, recv_counts_mpi, recv_displs_mpi, MPIU_INT, comm));
178: PetscCall(PetscMalloc4(size, &send_counts_crd, size, &send_displs_crd, size, &recv_counts_crd, size, &recv_displs_crd));
179: for (PetscMPIInt r = 0; r < size; ++r) {
180: PetscCall(PetscMPIIntCast((PetscInt64)send_counts_mpi[r] * dim, &send_counts_crd[r]));
181: PetscCall(PetscMPIIntCast((PetscInt64)send_displs_mpi[r] * dim, &send_displs_crd[r]));
182: PetscCall(PetscMPIIntCast((PetscInt64)recv_counts_mpi[r] * dim, &recv_counts_crd[r]));
183: PetscCall(PetscMPIIntCast((PetscInt64)recv_displs_mpi[r] * dim, &recv_displs_crd[r]));
184: }
185: PetscCallMPI(MPI_Alltoallv(send_crd, send_counts_crd, send_displs_crd, MPIU_REAL, recv_crd, recv_counts_crd, recv_displs_crd, MPIU_REAL, comm));
187: PetscCall(PetscFree4(send_counts_mpi, send_displs_mpi, recv_counts_mpi, recv_displs_mpi));
188: PetscCall(PetscFree4(send_counts_crd, send_displs_crd, recv_counts_crd, recv_displs_crd));
189: PetscCall(PetscFree2(send_counts, send_displs));
190: PetscCall(PetscFree2(recv_counts, recv_displs));
191: PetscCall(PetscFree(send_idx));
192: PetscCall(PetscFree(send_crd));
193: PetscCall(PetscFree(all_bboxes));
195: *n_obs_filt = total_recv;
196: *obs_idx_out = recv_idx;
197: *obs_coords_out = recv_crd;
198: PetscFunctionReturn(PETSC_SUCCESS);
199: }
201: /*
202: PetscDALETKFCoalesceNnzMinMax - In-place MAX-reduction of per-row nnz min/max across `comm`.
204: Coalesces (max, min) into a single MAX allreduce by negating the min, which halves the latency
205: versus two separate reductions. The sentinel `PETSC_INT_MAX` round-trips through negation, so a
206: rank with zero local rows does not pollute the global min; on return the sentinel is clamped to 0
207: to keep the value meaningful for viewers and downstream checks.
208: */
209: PETSC_INTERN PetscErrorCode PetscDALETKFCoalesceNnzMinMax(MPI_Comm comm, PetscInt *min_inout, PetscInt *max_inout)
210: {
211: PetscInt mm[2];
213: PetscFunctionBegin;
214: mm[0] = *max_inout;
215: mm[1] = -(*min_inout);
216: PetscCallMPI(MPIU_Allreduce(MPI_IN_PLACE, mm, 2, MPIU_INT, MPI_MAX, comm));
217: *max_inout = mm[0];
218: *min_inout = -mm[1];
219: if (*min_inout == PETSC_INT_MAX) *min_inout = 0;
220: PetscFunctionReturn(PETSC_SUCCESS);
221: }
223: /*
224: PetscDALETKFComputeObsCoords - Allocate and fill obs_locs[d] = H * xyz[d] for d in [0, dim).
226: Determines `dim` from the contiguous non-NULL prefix of `xyz` (with a contiguity check) and
227: returns a freshly allocated `Vec[dim]` whose entries are H-image vectors carrying coordinate
228: values at observation locations. Caller frees with `PetscDALETKFDestroyObsCoords()`.
229: */
230: PETSC_INTERN PetscErrorCode PetscDALETKFComputeObsCoords(Mat H, Vec xyz[], PetscInt *dim_out, Vec **obs_vecs_out)
231: {
232: MPI_Comm comm;
233: PetscInt dim = 0;
234: Vec *obs_vecs;
236: PetscFunctionBegin;
237: PetscCall(PetscObjectGetComm((PetscObject)H, &comm));
238: for (PetscInt d = 0; d < 3; ++d) {
239: if (xyz[d]) {
240: PetscCheck(d == dim, comm, PETSC_ERR_ARG_WRONG, "Coordinate slots must be contiguous from xyz[0]; got NULL before xyz[%" PetscInt_FMT "]", d);
241: dim++;
242: }
243: }
244: PetscCheck(dim >= 1, comm, PETSC_ERR_ARG_WRONG, "At least one coordinate vector required in xyz[0]");
245: PetscCall(PetscMalloc1(dim, &obs_vecs));
246: for (PetscInt d = 0; d < dim; ++d) {
247: PetscCall(MatCreateVecs(H, NULL, &obs_vecs[d]));
248: PetscCall(MatMult(H, xyz[d], obs_vecs[d]));
249: }
250: *dim_out = dim;
251: *obs_vecs_out = obs_vecs;
252: PetscFunctionReturn(PETSC_SUCCESS);
253: }
255: /*
256: PetscDALETKFDestroyObsCoords - Tear down the obs_vecs array allocated by PetscDALETKFComputeObsCoords().
257: */
258: PETSC_INTERN PetscErrorCode PetscDALETKFDestroyObsCoords(PetscInt dim, Vec **obs_vecs)
259: {
260: PetscFunctionBegin;
261: if (!*obs_vecs) PetscFunctionReturn(PETSC_SUCCESS);
262: for (PetscInt d = 0; d < dim; ++d) PetscCall(VecDestroy(&(*obs_vecs)[d]));
263: PetscCall(PetscFree(*obs_vecs));
264: *obs_vecs = NULL;
265: PetscFunctionReturn(PETSC_SUCCESS);
266: }
268: /*
269: PetscDALETKFAssembleQFromCSR - Materialize the localization Mat `Q` from a per-rank CSR triple.
271: Shared by the AIJ and Kokkos backends after each has produced (`row_counts`, `row_offsets`,
272: `col_indices`, `values`) describing every local row of `Q`. `mat_type` selects `MATAIJ` or
273: `MATAIJKOKKOS`. `H` is consulted only for its row ownership range, which determines the
274: diagonal-vs-off-diagonal split used for MPI preallocation. The output Mat is fully assembled.
275: */
276: PETSC_INTERN PetscErrorCode PetscDALETKFAssembleQFromCSR(Mat H, PetscInt n_vert_local, PetscInt n_obs_local, PetscInt n_obs_global, MatType mat_type, const PetscInt row_counts[], const PetscInt row_offsets[], const PetscInt col_indices[], const PetscScalar values[], Mat *Q)
277: {
278: MPI_Comm comm;
279: PetscInt rstart, cstart, cend;
280: PetscInt *d_nnz, *o_nnz;
281: PetscBool is_aij, is_aijkok, is_seqaij, is_mpiaij;
283: PetscFunctionBegin;
284: /* The seq/MPI AIJ preallocation calls below are no-ops for non-AIJ types, so an unsupported
285: mat_type would silently fall through to slow MatSetValues with no preallocation. Reject it
286: up front instead. */
287: PetscCall(PetscStrcmp(mat_type, MATAIJ, &is_aij));
288: PetscCall(PetscStrcmp(mat_type, MATAIJKOKKOS, &is_aijkok));
289: PetscCall(PetscStrcmp(mat_type, MATSEQAIJ, &is_seqaij));
290: PetscCall(PetscStrcmp(mat_type, MATMPIAIJ, &is_mpiaij));
291: PetscCheck(is_aij || is_aijkok || is_seqaij || is_mpiaij, PetscObjectComm((PetscObject)H), PETSC_ERR_SUP, "Unsupported Q mat_type \"%s\"; expected MATAIJ or MATAIJKOKKOS", mat_type);
292: PetscCall(PetscObjectGetComm((PetscObject)H, &comm));
293: PetscCall(MatGetOwnershipRange(H, &cstart, &cend));
295: PetscCall(MatCreate(comm, Q));
296: PetscCall(MatSetSizes(*Q, n_vert_local, n_obs_local, PETSC_DETERMINE, n_obs_global));
297: PetscCall(MatSetType(*Q, mat_type));
298: PetscCall(MatGetOwnershipRange(*Q, &rstart, NULL));
299: PetscCall(PetscCalloc1(n_vert_local, &d_nnz));
300: PetscCall(PetscCalloc1(n_vert_local, &o_nnz));
301: for (PetscInt i = 0; i < n_vert_local; ++i) {
302: PetscInt nnz = row_counts[i];
303: PetscInt off = row_offsets[i];
304: for (PetscInt k = 0; k < nnz; ++k) {
305: PetscInt col = col_indices[off + k];
306: if (col >= cstart && col < cend) d_nnz[i]++;
307: else o_nnz[i]++;
308: }
309: }
310: /* MatXAIJSetPreallocation dispatches to the right Seq/MPI variant for AIJ matrices. In serial
311: d_nnz already holds the full per-row count (cstart=0, cend=n_obs_global covers every column),
312: so the same dnnz/onnz pair is correct in both layouts. */
313: PetscCall(MatXAIJSetPreallocation(*Q, 1, d_nnz, o_nnz, NULL, NULL));
314: PetscCall(PetscFree(d_nnz));
315: PetscCall(PetscFree(o_nnz));
317: for (PetscInt i = 0; i < n_vert_local; ++i) {
318: PetscInt global_row = rstart + i;
319: PetscInt nnz = row_counts[i];
320: PetscInt off = row_offsets[i];
321: PetscCall(MatSetValues(*Q, 1, &global_row, nnz, &col_indices[off], &values[off], INSERT_VALUES));
322: }
323: PetscCall(MatAssemblyBegin(*Q, MAT_FINAL_ASSEMBLY));
324: PetscCall(MatAssemblyEnd(*Q, MAT_FINAL_ASSEMBLY));
325: PetscFunctionReturn(PETSC_SUCCESS);
326: }
328: /*
329: PetscDALETKFLogQStats - PetscInfo one-liner summarizing min/max nnz per row of `Q`.
330: */
331: PETSC_INTERN PetscErrorCode PetscDALETKFLogQStats(Mat Q, PetscDALETKFLocalizationType type, PetscReal radius, PetscInt n_vert_local, PetscInt n_obs_global, const PetscInt row_counts[])
332: {
333: MPI_Comm comm;
334: PetscInt local_min = PETSC_INT_MAX, local_max = 0;
336: PetscFunctionBegin;
337: PetscCall(PetscObjectGetComm((PetscObject)Q, &comm));
338: for (PetscInt i = 0; i < n_vert_local; ++i) {
339: if (row_counts[i] < local_min) local_min = row_counts[i];
340: if (row_counts[i] > local_max) local_max = row_counts[i];
341: }
342: PetscCall(PetscDALETKFCoalesceNnzMinMax(comm, &local_min, &local_max));
343: PetscCall(PetscInfo((PetscObject)Q, "LETKF localization (type=%s, radius=%g): %" PetscInt_FMT " vertices, %" PetscInt_FMT " obs, nnz/row min=%" PetscInt_FMT " max=%" PetscInt_FMT "\n", PetscDALETKFLocalizationTypes[type], (double)radius, n_vert_local, n_obs_global, local_min, local_max));
344: PetscFunctionReturn(PETSC_SUCCESS);
345: }
347: /*
348: PetscDALETKFCreateLocalizationMat_AIJ - host (`MATAIJ`) implementation of the localization-weight Mat `Q`.
350: Counterpart to `PetscDALETKFCreateLocalizationMat_Kokkos()`; same two-pass count/fill structure but plain C
351: with no Kokkos dependency. Selected by the `PetscDALETKFCreateLocalizationMat()` dispatcher when `H` is not
352: a Kokkos matrix type.
353: */
354: static PetscErrorCode PetscDALETKFCreateLocalizationMat_AIJ(PetscDALETKFLocalizationType type, PetscReal radius, Vec xyz[], PetscReal bd[], Mat H, Mat *Q, PetscInt *max_nnz_local, PetscInt *n_nnz_local)
355: {
356: PetscInt dim, n_vert_local, n_obs_global, n_obs_local, n_obs_cand;
357: PetscInt total_nnz = 0;
358: PetscInt64 total_nnz64 = 0;
359: PetscInt row_max = 0;
360: PetscInt *row_counts, *row_offsets, *col_indices, *obs_global_idx;
361: PetscScalar *values;
362: PetscReal cutoff, cutoff2;
363: PetscReal *vertex_coords, *obs_coords;
364: Vec *obs_vecs;
366: PetscFunctionBegin;
367: PetscCall(MatGetLocalSize(H, &n_obs_local, NULL));
368: PetscCall(MatGetSize(H, &n_obs_global, NULL));
369: PetscCall(PetscDALETKFComputeObsCoords(H, xyz, &dim, &obs_vecs));
370: PetscCall(VecGetLocalSize(xyz[0], &n_vert_local));
372: /* Vertex coordinates flattened as [vert][dim] (row-major). CSR row buffers are sized off
373: n_vert_local too, so allocate the trio in one shot. */
374: PetscCall(PetscMalloc3((size_t)n_vert_local * dim, &vertex_coords, n_vert_local, &row_counts, n_vert_local + 1, &row_offsets));
375: for (PetscInt d = 0; d < dim; ++d) {
376: const PetscScalar *local_coords_array;
377: PetscCall(VecGetArrayRead(xyz[d], &local_coords_array));
378: for (PetscInt i = 0; i < n_vert_local; ++i) vertex_coords[i * dim + d] = PetscRealPart(local_coords_array[i]);
379: PetscCall(VecRestoreArrayRead(xyz[d], &local_coords_array));
380: }
382: /* Single source of truth for the cutoff policy lives in letkf_kernels.h; LETKFCutoff() returns
383: the un-squared bound directly to avoid sqrt(r*r) round-trip FP error in the bbox prune
384: (matters for the BOXCAR kernel's strict (distance < radius) test on boundary obs). */
385: cutoff = LETKFCutoff(type, radius);
386: cutoff2 = cutoff * cutoff;
388: /* Bbox-pruned obs gather: replaces the all-to-all materialization of the global obs coord set. */
389: PetscCall(PetscDALETKFGatherObsBbox(dim, xyz, bd, cutoff, H, obs_vecs, &n_obs_cand, &obs_global_idx, &obs_coords));
391: /* Pass 1: Count nnz per row (only entries with positive kernel weight).
392: We pay the kernel evaluation twice (here and in Pass 2) to keep CSR allocations
393: sized exactly. The alternative -- allocate at the cutoff-bbox upper bound, fill,
394: then compact -- would peak at ~2x the memory and still touch every candidate.
395: Both passes go through LETKFRowWeight() (in letkf_kernels.h, shared with the
396: Kokkos backend) so the CPU and device kernels stay in lockstep. */
397: for (PetscInt i = 0; i < n_vert_local; ++i) {
398: PetscInt count = 0;
399: for (PetscInt j = 0; j < n_obs_cand; ++j) {
400: if (LETKFRowWeight(type, radius, cutoff2, dim, &vertex_coords[i * dim], &obs_coords[j * dim], bd) > 0.0) count++;
401: }
402: row_counts[i] = count;
403: }
405: /* Prefix sum for CSR row offsets, total nnz, and per-rank max nnz. Accumulate the running
406: total in 64-bit and cast so we trip a clear error instead of silently wrapping when
407: localization radius * obs density overflows PetscInt; mirrors the Kokkos backend
408: (kokkos/dalocalizationletkf.kokkos.cxx). The per-rank max feeds PetscDALETKFInstallQ()
409: without a downstream MatGetRow walk. */
410: row_offsets[0] = 0;
411: for (PetscInt i = 0; i < n_vert_local; ++i) {
412: total_nnz64 += (PetscInt64)row_counts[i];
413: row_offsets[i + 1] = row_offsets[i] + row_counts[i];
414: if (row_counts[i] > row_max) row_max = row_counts[i];
415: }
416: PetscCall(PetscIntCast(total_nnz64, &total_nnz));
418: /* Pass 2: Fill column indices and weights. The (w > 0.0) gate must match Pass 1 exactly
419: (same LETKFRowWeight() call) so each row writes precisely row_counts[i] entries.
420: Both passes feed identical operand sequences to LETKFRowWeight(), so the FP results are
421: bit-identical and the pass-1/pass-2 count match below holds in IEEE-754; the PetscAssert
422: is debug-only insurance against future drift in LETKFRowWeight itself. */
423: PetscCall(PetscMalloc2(total_nnz, &col_indices, total_nnz, &values));
424: for (PetscInt i = 0; i < n_vert_local; ++i) {
425: PetscInt offset = row_offsets[i];
426: PetscInt pos = 0;
427: for (PetscInt j = 0; j < n_obs_cand; ++j) {
428: PetscReal w = LETKFRowWeight(type, radius, cutoff2, dim, &vertex_coords[i * dim], &obs_coords[j * dim], bd);
429: if (w > 0.0) {
430: col_indices[offset + pos] = obs_global_idx[j];
431: values[offset + pos] = w;
432: pos++;
433: }
434: }
435: PetscAssert(pos == row_counts[i], PETSC_COMM_SELF, PETSC_ERR_PLIB, "LETKF localization Pass 1/2 mismatch on row %" PetscInt_FMT ": pass1=%" PetscInt_FMT " pass2=%" PetscInt_FMT, i, row_counts[i], pos);
436: }
438: PetscCall(PetscDALETKFAssembleQFromCSR(H, n_vert_local, n_obs_local, n_obs_global, MATAIJ, row_counts, row_offsets, col_indices, values, Q));
439: PetscCall(PetscDALETKFLogQStats(*Q, type, radius, n_vert_local, n_obs_global, row_counts));
441: *max_nnz_local = row_max;
442: *n_nnz_local = total_nnz;
443: PetscCall(PetscDALETKFDestroyObsCoords(dim, &obs_vecs));
444: PetscCall(PetscFree3(vertex_coords, row_counts, row_offsets));
445: PetscCall(PetscFree2(col_indices, values));
446: PetscCall(PetscFree(obs_coords));
447: PetscCall(PetscFree(obs_global_idx));
448: PetscFunctionReturn(PETSC_SUCCESS);
449: }
451: /*
452: PetscDALETKFCreateLocalizationMat - construct the LETKF localization-weight Mat `Q` for a built-in
453: distance kernel. Validates the common arguments and dispatches to either the Kokkos backend or the
454: host `MATAIJ` backend based on `use_kokkos`. The caller chooses the backend so this matches the
455: analysis-time backend selection (which keys off the obs-error covariance Mat `da->R`); using `H`'s
456: type here would dispatch inconsistently when the user has a Kokkos `H` but a CPU `R` (or vice versa).
458: Output `Q` has rows indexed by local grid vertices and columns indexed by global observations; the
459: caller owns the returned Mat.
460: */
461: PETSC_INTERN PetscErrorCode PetscDALETKFCreateLocalizationMat(PetscDALETKFLocalizationType type, PetscReal radius, Vec xyz[], PetscReal bd[], Mat H, PetscBool use_kokkos, Mat *Q, PetscInt *max_nnz_local, PetscInt *n_nnz_local)
462: {
463: MPI_Comm comm;
465: PetscFunctionBegin;
466: PetscAssertPointer(xyz, 3);
467: PetscAssertPointer(bd, 4);
469: PetscAssertPointer(Q, 7);
470: PetscAssertPointer(max_nnz_local, 8);
471: PetscAssertPointer(n_nnz_local, 9);
472: PetscCall(PetscObjectGetComm((PetscObject)H, &comm));
473: PetscCheck(type == PETSCDA_LETKF_LOC_GASPARI_COHN || type == PETSCDA_LETKF_LOC_GAUSSIAN || type == PETSCDA_LETKF_LOC_BOXCAR, comm, PETSC_ERR_ARG_WRONG, "Built-in kernel required, got localization type %d", (int)type);
474: PetscCheck(radius > 0, comm, PETSC_ERR_ARG_OUTOFRANGE, "Localization radius must be positive, got %g", (double)radius);
475: /* xyz[] contiguity and dim>=1 are enforced by PetscDALETKFComputeObsCoords() inside the backends. */
476: #if PetscDefined(HAVE_KOKKOS_KERNELS) && !PetscDefined(USE_COMPLEX)
477: if (use_kokkos) {
478: PetscCall(PetscDALETKFCreateLocalizationMat_Kokkos(type, radius, xyz, bd, H, Q, max_nnz_local, n_nnz_local));
479: PetscFunctionReturn(PETSC_SUCCESS);
480: }
481: #endif
482: PetscCall(PetscDALETKFCreateLocalizationMat_AIJ(type, radius, xyz, bd, H, Q, max_nnz_local, n_nnz_local));
483: PetscFunctionReturn(PETSC_SUCCESS);
484: }