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