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