Actual source code: dalocalizationletkf.kokkos.cxx
1: #include <petsc.h>
2: #include <petscmat.h>
3: #include <petsc_kokkos.hpp>
4: #include <cmath>
5: #include <Kokkos_Core.hpp>
6: #include <../src/ml/da/impls/ensemble/letkf/letkf.h>
7: #include <../src/ml/da/impls/ensemble/letkf/letkf_kernels.h>
9: /*
10: PetscDALETKFCreateLocalizationMat_Kokkos - Kokkos (`MATAIJKOKKOS`) implementation of the localization
11: weight Mat `Q`.
13: Selected by the `PetscDALETKFCreateLocalizationMat()` dispatcher when the caller requests the Kokkos
14: backend (i.e. when `da->R` has a Kokkos matrix type). The host counterpart
15: `PetscDALETKFCreateLocalizationMat_AIJ()` produces a numerically identical Q for the polynomial kernels
16: (`PETSCDA_LETKF_LOC_GASPARI_COHN`, `PETSCDA_LETKF_LOC_BOXCAR`) and a bit-comparable Q for
17: `PETSCDA_LETKF_LOC_GAUSSIAN` modulo `exp()` rounding.
18: */
19: PETSC_INTERN PetscErrorCode PetscDALETKFCreateLocalizationMat_Kokkos(PetscDALETKFLocalizationType type, PetscReal radius, Vec xyz[], PetscReal bd[], Mat H, Mat *Q, PetscInt *max_nnz_local, PetscInt *n_nnz_local)
20: {
21: using ExecSpace = Kokkos::DefaultExecutionSpace;
22: using MemSpace = ExecSpace::memory_space;
23: using DevScalar2D = Kokkos::View<PetscScalar **, Kokkos::LayoutLeft, MemSpace>;
24: using HostScalar2D = Kokkos::View<PetscScalar **, Kokkos::LayoutLeft, Kokkos::HostSpace>;
25: using DevInt1D = Kokkos::View<PetscInt *, Kokkos::LayoutLeft, MemSpace>;
26: using DevScalar1D = Kokkos::View<PetscScalar *, Kokkos::LayoutLeft, MemSpace>;
27: using HostInt1D = Kokkos::View<PetscInt *, Kokkos::LayoutLeft, Kokkos::HostSpace>;
28: using HostScalar1D = Kokkos::View<PetscScalar *, Kokkos::LayoutLeft, Kokkos::HostSpace>;
29: PetscInt dim, n_vert_local, n_obs_global, n_obs_local, n_obs_cand;
30: PetscInt total_nnz = 0;
31: PetscInt64 total_nnz64 = 0;
32: PetscInt row_max = 0;
33: PetscInt *obs_global_idx_host;
34: PetscReal cutoff, cutoff2;
35: PetscReal *obs_coords_host_buf;
36: Vec *obs_vecs;
37: DevScalar2D vertex_coords_dev;
38: HostScalar2D vertex_coords_host;
39: Kokkos::View<PetscReal **, Kokkos::LayoutRight, MemSpace> obs_coords_dev;
40: Kokkos::View<PetscInt *, MemSpace> obs_global_idx_dev;
41: Kokkos::View<PetscReal *, MemSpace> bd_dev;
42: Kokkos::View<PetscReal *, Kokkos::HostSpace> bd_host;
43: DevInt1D row_counts_dev, row_offsets_dev, col_indices_dev;
44: DevScalar1D values_dev;
45: HostInt1D row_counts_host, row_offsets_host, col_indices_host;
46: HostScalar1D values_host;
47: #if PetscDefined(USE_DEBUG)
48: DevInt1D actual_counts_dev;
49: HostInt1D actual_counts_host;
50: #endif
52: PetscFunctionBegin;
53: /* Single source of truth for the cutoff policy lives in letkf_kernels.h; see CPU
54: dalocalizationletkf.c for why this is computed directly rather than via sqrt(cutoff^2). */
55: cutoff = LETKFCutoff(type, radius);
56: cutoff2 = cutoff * cutoff;
57: PetscCall(PetscKokkosInitializeCheck());
58: PetscCall(MatGetLocalSize(H, &n_obs_local, NULL));
59: PetscCall(MatGetSize(H, &n_obs_global, NULL));
60: PetscCall(PetscDALETKFComputeObsCoords(H, xyz, &dim, &obs_vecs));
61: PetscCall(VecGetLocalSize(xyz[0], &n_vert_local));
63: /* Copy vertex coordinates to device */
64: PetscCallCXX(vertex_coords_dev = Kokkos::View<PetscScalar **, Kokkos::LayoutLeft, MemSpace>("vertex_coords", n_vert_local, dim));
65: PetscCallCXX(vertex_coords_host = Kokkos::View<PetscScalar **, Kokkos::LayoutLeft, Kokkos::HostSpace>("vertex_coords_host", n_vert_local, dim));
66: for (PetscInt d = 0; d < dim; ++d) {
67: const PetscScalar *local_coords_array;
68: PetscCall(VecGetArrayRead(xyz[d], &local_coords_array));
69: for (PetscInt i = 0; i < n_vert_local; ++i) vertex_coords_host(i, d) = local_coords_array[i];
70: PetscCall(VecRestoreArrayRead(xyz[d], &local_coords_array));
71: }
72: PetscCallCXX(Kokkos::deep_copy(vertex_coords_dev, vertex_coords_host));
74: /* Bbox-pruned obs gather: each rank receives only obs whose location can fall within the kernel
75: cutoff of any vertex it owns. Keeps obs_coords device memory bounded by the local working set
76: instead of the global obs count. Output buffers are PetscMalloc1'd; we wrap them in Kokkos
77: unmanaged host views and deep_copy onto the device. */
78: PetscCall(PetscDALETKFGatherObsBbox(dim, xyz, bd, cutoff, H, obs_vecs, &n_obs_cand, &obs_global_idx_host, &obs_coords_host_buf));
80: PetscCallCXX(obs_coords_dev = Kokkos::View<PetscReal **, Kokkos::LayoutRight, MemSpace>("obs_coords", n_obs_cand, dim));
81: PetscCallCXX(obs_global_idx_dev = Kokkos::View<PetscInt *, MemSpace>("obs_global_idx", n_obs_cand));
82: PetscCallCXX(Kokkos::deep_copy(obs_coords_dev, Kokkos::View<const PetscReal **, Kokkos::LayoutRight, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>>(obs_coords_host_buf, n_obs_cand, dim)));
83: PetscCallCXX(Kokkos::deep_copy(obs_global_idx_dev, Kokkos::View<const PetscInt *, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>>(obs_global_idx_host, n_obs_cand)));
85: /* Copy boundary data to device */
86: PetscCallCXX(bd_dev = Kokkos::View<PetscReal *, MemSpace>("bd_dev", dim));
87: PetscCallCXX(bd_host = Kokkos::View<PetscReal *, Kokkos::HostSpace>("bd_host", dim));
88: for (PetscInt d = 0; d < dim; ++d) bd_host(d) = bd[d];
89: PetscCallCXX(Kokkos::deep_copy(bd_dev, bd_host));
91: /* Pass 1: Count nnz per row (only entries with positive kernel weight).
92: Same trade-off as the host backend: re-evaluate the kernel in Pass 2 rather than
93: allocate and compact at the cutoff-bbox upper bound. The arithmetic is cheap on the
94: device; the global memory writes saved by exact sizing dominate. */
95: PetscCallCXX(row_counts_dev = Kokkos::View<PetscInt *, Kokkos::LayoutLeft, MemSpace>("row_counts", n_vert_local));
97: Kokkos::parallel_for(
98: "CountNnzPerRow", Kokkos::RangePolicy<ExecSpace>(0, n_vert_local), KOKKOS_LAMBDA(const PetscInt i) {
99: PetscReal v_coords[3] = {0.0, 0.0, 0.0};
100: for (PetscInt dd = 0; dd < dim; ++dd) v_coords[dd] = PetscRealPart(vertex_coords_dev(i, dd));
102: PetscInt count = 0;
103: for (PetscInt j = 0; j < n_obs_cand; ++j) {
104: if (LETKFRowWeight(type, radius, cutoff2, dim, v_coords, &obs_coords_dev(j, 0), bd_dev.data()) > 0.0) count++;
105: }
106: row_counts_dev(i) = count;
107: });
108: Kokkos::fence();
110: /* Copy row counts to host for preallocation */
111: PetscCallCXX(row_counts_host = Kokkos::View<PetscInt *, Kokkos::LayoutLeft, Kokkos::HostSpace>("row_counts_host", n_vert_local));
112: PetscCallCXX(Kokkos::deep_copy(row_counts_host, row_counts_dev));
114: /* Compute prefix sum on host for CSR row offsets */
115: PetscCallCXX(row_offsets_dev = Kokkos::View<PetscInt *, Kokkos::LayoutLeft, MemSpace>("row_offsets", n_vert_local + 1));
116: PetscCallCXX(row_offsets_host = Kokkos::View<PetscInt *, Kokkos::LayoutLeft, Kokkos::HostSpace>("row_offsets_host", n_vert_local + 1));
117: row_offsets_host(0) = 0;
118: for (PetscInt i = 0; i < n_vert_local; ++i) row_offsets_host(i + 1) = row_offsets_host(i) + row_counts_host(i);
119: PetscCallCXX(Kokkos::deep_copy(row_offsets_dev, row_offsets_host));
121: /* Total nnz for output arrays and per-rank max nnz for downstream sizing. Accumulate in
122: 64-bit and cast so we trip a clear error instead of silently wrapping when localization
123: radius * obs density overflows PetscInt. The per-rank max feeds PetscDALETKFInstallQ()
124: without a downstream MatGetRow walk. */
125: for (PetscInt i = 0; i < n_vert_local; ++i) {
126: total_nnz64 += (PetscInt64)row_counts_host(i);
127: if (row_counts_host(i) > row_max) row_max = row_counts_host(i);
128: }
129: PetscCall(PetscIntCast(total_nnz64, &total_nnz));
131: /* Pass 2: Fill column indices and weights */
132: PetscCallCXX(col_indices_dev = Kokkos::View<PetscInt *, Kokkos::LayoutLeft, MemSpace>("col_indices", total_nnz));
133: PetscCallCXX(values_dev = Kokkos::View<PetscScalar *, Kokkos::LayoutLeft, MemSpace>("values", total_nnz));
135: #if PetscDefined(USE_DEBUG)
136: /* Mirror the CPU backend's PetscAssert(pos == row_counts[i]): record the actual Pass-2 write
137: count per row in a device View and verify on host that it matches Pass 1. Allocated only in
138: debug builds; release builds skip both the View and the per-row write. */
139: PetscCallCXX(actual_counts_dev = Kokkos::View<PetscInt *, Kokkos::LayoutLeft, MemSpace>("actual_counts", n_vert_local));
140: #endif
142: Kokkos::parallel_for(
143: "FillLocalizationMatrix", Kokkos::RangePolicy<ExecSpace>(0, n_vert_local), KOKKOS_LAMBDA(const PetscInt i) {
144: PetscInt offset = row_offsets_dev(i);
146: PetscReal v_coords[3] = {0.0, 0.0, 0.0};
147: for (PetscInt dd = 0; dd < dim; ++dd) v_coords[dd] = PetscRealPart(vertex_coords_dev(i, dd));
149: PetscInt pos = 0;
150: for (PetscInt j = 0; j < n_obs_cand; ++j) {
151: PetscReal w = LETKFRowWeight(type, radius, cutoff2, dim, v_coords, &obs_coords_dev(j, 0), bd_dev.data());
152: if (w > 0.0) {
153: col_indices_dev(offset + pos) = obs_global_idx_dev(j);
154: values_dev(offset + pos) = w;
155: pos++;
156: }
157: }
158: #if PetscDefined(USE_DEBUG)
159: actual_counts_dev(i) = pos;
160: #endif
161: });
162: Kokkos::fence();
164: #if PetscDefined(USE_DEBUG)
165: PetscCallCXX(actual_counts_host = Kokkos::View<PetscInt *, Kokkos::LayoutLeft, Kokkos::HostSpace>("actual_counts_host", n_vert_local));
166: PetscCallCXX(Kokkos::deep_copy(actual_counts_host, actual_counts_dev));
167: for (PetscInt i = 0; i < n_vert_local; ++i)
168: PetscCheck(actual_counts_host(i) == row_counts_host(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_host(i), actual_counts_host(i));
169: #endif
171: /* Copy results to host */
172: PetscCallCXX(col_indices_host = Kokkos::View<PetscInt *, Kokkos::LayoutLeft, Kokkos::HostSpace>("col_indices_host", total_nnz));
173: PetscCallCXX(values_host = Kokkos::View<PetscScalar *, Kokkos::LayoutLeft, Kokkos::HostSpace>("values_host", total_nnz));
174: PetscCallCXX(Kokkos::deep_copy(col_indices_host, col_indices_dev));
175: PetscCallCXX(Kokkos::deep_copy(values_host, values_dev));
177: PetscCall(PetscDALETKFAssembleQFromCSR(H, n_vert_local, n_obs_local, n_obs_global, MATAIJKOKKOS, row_counts_host.data(), row_offsets_host.data(), col_indices_host.data(), values_host.data(), Q));
178: PetscCall(PetscDALETKFLogQStats(*Q, type, radius, n_vert_local, n_obs_global, row_counts_host.data()));
180: *max_nnz_local = row_max;
181: *n_nnz_local = total_nnz;
182: PetscCall(PetscDALETKFDestroyObsCoords(dim, &obs_vecs));
183: PetscCall(PetscFree(obs_global_idx_host));
184: PetscCall(PetscFree(obs_coords_host_buf));
185: PetscFunctionReturn(PETSC_SUCCESS);
186: }