Actual source code: letkf_local_analysis.kokkos.cxx

  1: #include <../src/ml/da/impls/ensemble/letkf/letkf.h>
  2: #include <Kokkos_Core.hpp>
  3: #include <KokkosBlas.hpp>
  4: #include <climits>

  6: #if defined(KOKKOS_ENABLE_CUDA)
  7:   #include <cusolverDn.h>
  8:   #include <cuda_runtime.h>
  9: #include <petscdevice_cuda.h>
 10: #elif defined(KOKKOS_ENABLE_HIP)
 11:   #include <rocsolver/rocsolver.h>
 12:   #include <hip/hip_runtime.h>
 13: #include <petscdevice_hip.h>
 14: #elif defined(KOKKOS_ENABLE_SYCL)
 15:   #include <oneapi/mkl.hpp>
 16:   #include <sycl/sycl.hpp>
 17: #endif

 19: /* Shared device-View aliases used throughout the BatchedEigenSolve* dispatch chain and
 20:    by the per-mirror device locals in PetscDALETKFLocalAnalysis_Kokkos. */
 21: using LETKFExecSpace = Kokkos::DefaultExecutionSpace;
 22: using LETKFView3D    = Kokkos::View<PetscScalar ***, Kokkos::LayoutLeft, LETKFExecSpace>;
 23: using LETKFView2D    = Kokkos::View<PetscScalar **, Kokkos::LayoutLeft, LETKFExecSpace>;
 24: using LETKFView1D    = Kokkos::View<PetscScalar *, Kokkos::LayoutLeft, LETKFExecSpace>;

 26: /* Floor used when dividing by Lambda_i from the per-vertex eigendecomposition of
 27:    T = (1/rho)*I + S^T*S. T is SPD by construction so eigenvalues are positive in exact arithmetic,
 28:    but device eigensolvers (cusolver syevj, rocsolver syevd, oneMKL syevd) can produce tiny
 29:    rounding-magnitude eigenvalues on near-degenerate spectra; without a floor, dividing by them
 30:    produces inf/nan in temp2 and inv_sqrt_lambda_i and silently corrupts the analysis in release
 31:    builds (the DEBUG CheckLambda block only flags eigenvalues below -1e-8, which lets near-zero
 32:    positives pass through). Scaled by PETSC_MACHINE_EPSILON so the floor follows precision; the
 33:    absolute floor matches the original 1.0e-14 in double precision and adapts down/up for
 34:    single/quad. */
 35: static constexpr PetscReal LETKF_EIGEN_EPS = (PetscReal)100.0 * PETSC_MACHINE_EPSILON;

 37: /* ========================================================================== */
 38: /*                    Batched Eigendecomposition for LETKF                    */
 39: /* ========================================================================== */

 41: /* Structure to hold reusable workspace for eigensolvers.
 42:    Lifecycle is manual: allocated in PetscDALETKFLocalAnalysis_Kokkos/GlobalAnalysis_Kokkos when
 43:    first needed and freed in PetscDALETKFDestroyLocalization_Kokkos. We do not provide a
 44:    destructor because the cleanup uses CUDA/HIP/SYCL APIs that need PETSc error checking
 45:    (PetscCallCUDA/HIP, sycl::free with a queue), which cannot be expressed inside a C++
 46:    destructor body that is supposed to be noexcept. */
 47: struct EigenWorkspace {
 48:   /* Tracking for reuse */
 49:   PetscInt max_chunk_size;
 50:   PetscInt m;
 51:   PetscInt max_nnz;

 53:   /* Persistent Kokkos Views */
 54:   using exec_space = Kokkos::DefaultExecutionSpace;
 55:   using view_3d    = Kokkos::View<PetscScalar ***, Kokkos::LayoutLeft, exec_space>;
 56:   using view_2d    = Kokkos::View<PetscScalar **, Kokkos::LayoutLeft, exec_space>;

 58:   view_3d S_batch;
 59:   view_3d T_batch;
 60:   view_3d V_batch;
 61:   view_2d Lambda_batch;
 62:   view_3d T_sqrt_batch;
 63:   view_2d w_batch;
 64:   view_2d delta_batch;
 65:   view_2d y_batch;
 66:   view_2d y_mean_batch;
 67:   view_2d r_inv_sqrt_batch;
 68:   view_2d temp1_batch;
 69:   view_2d temp2_batch;
 70:   view_2d inv_sqrt_lambda_batch;

 72:   /* Host workspace */
 73:   PetscScalar *all_v;
 74:   PetscReal   *all_lambda;
 75:   PetscScalar *all_work;
 76: #if PetscDefined(USE_COMPLEX)
 77:   PetscReal *all_rwork;
 78: #endif
 79:   PetscBLASInt lwork;
 80:   PetscBLASInt n_blas;

 82:   /* Device workspace */
 83: #if defined(KOKKOS_ENABLE_CUDA) || defined(KOKKOS_ENABLE_HIP) || defined(KOKKOS_ENABLE_SYCL)
 84:   #if defined(KOKKOS_ENABLE_CUDA)
 85:   syevjInfo_t  syevj_params;
 86:   PetscScalar *d_work;
 87:   int         *d_info;
 88:   PetscScalar *d_A_contig;
 89:   PetscScalar *d_W_contig;
 90:   int          lwork_device;
 91:   #elif defined(KOKKOS_ENABLE_HIP)
 92:   PetscScalar *d_work;
 93:   int         *d_info;
 94:   PetscScalar *d_A_contig;
 95:   PetscScalar *d_W_contig;
 96:   int          lwork_device;
 97:   #elif defined(KOKKOS_ENABLE_SYCL)
 98:   PetscScalar *d_work;
 99:   int         *d_info;
100:   PetscScalar *d_A_contig;
101:   PetscScalar *d_W_contig;
102:   int          lwork_device;
103:   #endif
104: #endif

106:   EigenWorkspace() : max_chunk_size(0), m(0), max_nnz(0), all_v(nullptr), all_lambda(nullptr), all_work(nullptr)
107:   {
108: #if PetscDefined(USE_COMPLEX)
109:     all_rwork = nullptr;
110: #endif
111: #if defined(KOKKOS_ENABLE_CUDA)
112:     d_work       = nullptr;
113:     d_info       = nullptr;
114:     d_A_contig   = nullptr;
115:     d_W_contig   = nullptr;
116:     syevj_params = nullptr;
117: #elif defined(KOKKOS_ENABLE_HIP) || defined(KOKKOS_ENABLE_SYCL)
118:     d_work     = nullptr;
119:     d_info     = nullptr;
120:     d_A_contig = nullptr;
121:     d_W_contig = nullptr;
122: #endif
123:   }
124: };

126: /*
127:   BatchedEigenSolve_Host - Compute eigendecomposition for a batch of symmetric matrices (CPU version)

129:   Input Parameters:
130: + T_batch      - batch of symmetric matrices (n_batch x n_size x n_size)
131: . n_batch      - number of matrices in the batch
132: - n_size       - size of each matrix (m x m)
133: - work         - reusable workspace structure

135:   Output Parameters:
136: + Lambda_batch - eigenvalues for each matrix (n_batch x n_size)
137: - V_batch      - eigenvectors for each matrix (n_batch x n_size x n_size)

139:   Notes:
140:   Uses LAPACK's syev routine to compute eigendecomposition sequentially on host.
141: */
142: #if !defined(KOKKOS_ENABLE_CUDA) && !defined(KOKKOS_ENABLE_HIP) && !defined(KOKKOS_ENABLE_SYCL)
143: #include <petscblaslapack.h>
144: static PetscErrorCode BatchedEigenSolve_Host(LETKFView3D T_batch, LETKFView2D Lambda_batch, LETKFView3D V_batch, PetscInt n_batch, PetscInt n_size, EigenWorkspace *work)
145: {
146:   PetscFunctionBegin;
147:   /* In the host-only build path the batch views already live in HostSpace, so LAPACK can
148:      read T_batch and write Lambda_batch/V_batch directly. The mirror+deep_copy round-trip
149:      used in the device path would be a no-op here, and some Kokkos+Serial configurations do
150:      not expose ::HostMirror on a View parameterized with an exec-space tag. */

152:   /* Use pre-allocated workspace */
153:   PetscScalar *all_v      = work->all_v;
154:   PetscReal   *all_lambda = work->all_lambda;
155:   PetscScalar *all_work   = work->all_work;
156:   PetscBLASInt lwork      = work->lwork;
157:   PetscBLASInt n_blas     = work->n_blas;
158:   #if PetscDefined(USE_COMPLEX)
159:   PetscReal *all_rwork = work->all_rwork;
160:   #endif

162:   /* Process each matrix in parallel on host using LAPACK */
163:   Kokkos::parallel_for(
164:     "BatchedEigenSolve_Host", Kokkos::RangePolicy<Kokkos::DefaultHostExecutionSpace>(0, n_batch), KOKKOS_LAMBDA(const int i) {
165:       PetscBLASInt n   = n_blas;
166:       PetscBLASInt lda = n;
167:       PetscBLASInt info;
168:       PetscBLASInt lw = lwork;

170:       /* Pointers for this matrix */
171:       PetscScalar *v_ptr      = all_v + i * n_size * n_size;
172:       PetscReal   *lambda_ptr = all_lambda + i * n_size;
173:       PetscScalar *work_ptr   = all_work + i * lwork;
174:   #if PetscDefined(USE_COMPLEX)
175:       PetscReal *rwork_ptr = all_rwork + i * (3 * n_size - 2);
176:   #endif

178:       /* Copy T_batch(i, :, :) to v_ptr (column-major) */
179:       for (PetscInt j = 0; j < n_size; j++) {
180:         for (PetscInt k = 0; k < n_size; k++) v_ptr[k + j * n_size] = T_batch(i, k, j);
181:       }

183:     /* Compute eigendecomposition: T = V * Lambda * V^T */
184:   #if PetscDefined(USE_COMPLEX)
185:       LAPACKsyev_("V", "U", &n, v_ptr, &lda, lambda_ptr, work_ptr, &lw, rwork_ptr, &info);
186:   #else
187:       LAPACKsyev_("V", "U", &n, v_ptr, &lda, lambda_ptr, work_ptr, &lw, &info);
188:   #endif

190:       /* Kokkos::parallel_for() cannot return error codes; abort the parallel region instead. */
191:       if (info != 0) Kokkos::abort("LAPACK eigendecomposition failed in parallel region");

193:       /* Write results directly back to the batch views (already in HostSpace) */
194:       for (PetscInt j = 0; j < n_size; j++) {
195:         Lambda_batch(i, j) = (PetscScalar)lambda_ptr[j];
196:         for (PetscInt k = 0; k < n_size; k++) V_batch(i, k, j) = v_ptr[k + j * n_size];
197:       }
198:     });
199:   PetscFunctionReturn(PETSC_SUCCESS);
200: }
201: #endif

203: /*
204:   BatchedEigenSolve_Device - Compute eigendecomposition for a batch of symmetric matrices (Device version)

206:   Input Parameters:
207: + T_batch      - batch of symmetric matrices (n_batch x n_size x n_size)
208: . n_batch      - number of matrices in the batch
209: - n_size       - size of each matrix (m x m)
210: - device_handle - device-specific solver handle (cusolverDnHandle_t, rocblas_handle, or sycl::queue*)
211: - work         - reusable workspace structure

213:   Output Parameters:
214: + Lambda_batch - eigenvalues for each matrix (n_batch x n_size)
215: - V_batch      - eigenvectors for each matrix (n_batch x n_size x n_size)

217:   Notes:
218:   Uses vendor-specific batched symmetric eigensolvers:
219:   - CUDA: cuSOLVER's syevjBatched
220:   - HIP: rocSOLVER's rocsolver_dsyevj_batched
221:   - SYCL: oneMKL's syevd_batch
222: */
223: #if defined(KOKKOS_ENABLE_CUDA) || defined(KOKKOS_ENABLE_HIP) || defined(KOKKOS_ENABLE_SYCL)
224:   #if defined(KOKKOS_ENABLE_CUDA)
225: static PetscErrorCode BatchedEigenSolve_Device(LETKFView3D T_batch, LETKFView2D Lambda_batch, LETKFView3D V_batch, PetscInt n_batch, PetscInt n_size, cusolverDnHandle_t cusolverH, EigenWorkspace *work)
226: {
227:   PetscFunctionBegin;
228:     #if PetscDefined(USE_COMPLEX)
229:   /* cuSOLVER's *syevjBatched is real-only (Ssyevj/Dsyevj); under complex the call would type-error.
230:      The dispatcher gates the Kokkos path off when PETSC_USE_COMPLEX is set, so this is unreachable
231:      in practice; SETERRQ here as defense-in-depth in case that gate ever changes. */
232:   SETERRQ(PETSC_COMM_SELF, PETSC_ERR_SUP, "Complex numbers not supported on CUDA backend for LETKF");
233:     #else
234:   cusolverStatus_t cusolver_status;
235:   syevjInfo_t      syevj_params = work->syevj_params;
236:   PetscScalar     *d_work       = work->d_work;
237:   int             *d_info       = work->d_info;
238:   PetscScalar     *d_A_contig   = work->d_A_contig;
239:   PetscScalar     *d_W_contig   = work->d_W_contig;
240:   int              lwork        = work->lwork_device;
241:   int             *h_info       = nullptr;
242:   /* Copy T_batch to contiguous layout for cuSOLVER */
243:   Kokkos::parallel_for(
244:     "ReorganizeForCuSOLVER", Kokkos::RangePolicy<Kokkos::DefaultExecutionSpace>(0, n_batch), KOKKOS_LAMBDA(const int i) {
245:       for (int j = 0; j < n_size; j++) {
246:         for (int k = 0; k < n_size; k++) d_A_contig[i * n_size * n_size + k * n_size + j] = T_batch(i, j, k);
247:       }
248:     });
249:   Kokkos::fence();

251:       /* Solve batched eigendecomposition */
252:       #if PetscDefined(USE_REAL_SINGLE)
253:   cusolver_status = cusolverDnSsyevjBatched(cusolverH, CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_UPPER, n_size, d_A_contig, n_size, d_W_contig, d_work, lwork, d_info, syevj_params, n_batch);
254:       #else
255:   cusolver_status = cusolverDnDsyevjBatched(cusolverH, CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_UPPER, n_size, d_A_contig, n_size, d_W_contig, d_work, lwork, d_info, syevj_params, n_batch);
256:       #endif
257:   PetscCheck(cusolver_status == CUSOLVER_STATUS_SUCCESS, PETSC_COMM_SELF, PETSC_ERR_LIB, "cusolverDn*syevjBatched failed");

259:   /* Check info */
260:   PetscCall(PetscMalloc1(n_batch, &h_info));
261:   PetscCallCUDA(cudaMemcpy(h_info, d_info, sizeof(int) * n_batch, cudaMemcpyDeviceToHost));
262:   for (PetscInt i = 0; i < n_batch; i++) PetscCheck(h_info[i] == 0, PETSC_COMM_SELF, PETSC_ERR_LIB, "cuSOLVER eigendecomposition failed for matrix %" PetscInt_FMT ": info=%d", i, h_info[i]);
263:   PetscCall(PetscFree(h_info));

265:   /* Copy results back from contiguous layout to V_batch */
266:   Kokkos::parallel_for(
267:     "CopyResultsBack", Kokkos::RangePolicy<Kokkos::DefaultExecutionSpace>(0, n_batch), KOKKOS_LAMBDA(const int i) {
268:       for (int j = 0; j < n_size; j++) {
269:         for (int k = 0; k < n_size; k++) V_batch(i, j, k) = d_A_contig[i * n_size * n_size + k * n_size + j];
270:         /* CUDA-12.6 nvcc compiler hangs if this line is placed before the V_batch loop. */
271:         Lambda_batch(i, j) = d_W_contig[i * n_size + j];
272:       }
273:     });
274:   Kokkos::fence();
275:     #endif
276:   PetscFunctionReturn(PETSC_SUCCESS);
277: }
278:   #elif defined(KOKKOS_ENABLE_HIP)
279: static PetscErrorCode BatchedEigenSolve_Device(LETKFView3D T_batch, LETKFView2D Lambda_batch, LETKFView3D V_batch, PetscInt n_batch, PetscInt n_size, rocblas_handle rocblasH, EigenWorkspace *work)
280: {
281:   PetscFunctionBegin;
282:     #if PetscDefined(USE_COMPLEX)
283:   /* Bail out before any kernel launch: the workspace setup leaves d_A_contig/d_W_contig/d_work/d_info
284:      as nullptr in complex mode (rocsolver_*syevd has no complex variant we wrap), so the
285:      ReorganizeForRocSOLVER parallel_for below would do a null device write before this error fired. */
286:   SETERRQ(PETSC_COMM_SELF, PETSC_ERR_SUP, "Complex numbers not supported on HIP backend for LETKF");
287:     #else
288:   PetscScalar *d_work     = work->d_work;
289:   int         *d_info     = work->d_info;
290:   PetscScalar *d_A_contig = work->d_A_contig;
291:   PetscScalar *d_W_contig = work->d_W_contig;
292:   int         *h_info     = nullptr;

294:   /* Copy T_batch to contiguous layout for rocSOLVER */
295:   Kokkos::parallel_for(
296:     "ReorganizeForRocSOLVER", Kokkos::RangePolicy<Kokkos::DefaultExecutionSpace>(0, n_batch), KOKKOS_LAMBDA(const int i) {
297:       for (int j = 0; j < n_size; j++) {
298:         for (int k = 0; k < n_size; k++) d_A_contig[i * n_size * n_size + k * n_size + j] = T_batch(i, j, k);
299:       }
300:     });
301:   Kokkos::fence();

303:   /* rocSOLVER doesn't have a native batched syevj, so we loop over batch.
304:      Use rocsolver_*syevd which is more efficient than calling syev in a loop. */
305:   for (int i = 0; i < n_batch; i++) {
306:     PetscScalar   *A_ptr    = d_A_contig + i * n_size * n_size;
307:     PetscScalar   *W_ptr    = d_W_contig + i * n_size;
308:     int           *info_ptr = d_info + i;
309:     rocblas_status hip_status;

311:       #if PetscDefined(USE_REAL_SINGLE)
312:     hip_status = rocsolver_ssyevd(rocblasH, rocblas_evect_original, rocblas_fill_upper, n_size, A_ptr, n_size, W_ptr, d_work, info_ptr);
313:       #else
314:     hip_status = rocsolver_dsyevd(rocblasH, rocblas_evect_original, rocblas_fill_upper, n_size, A_ptr, n_size, W_ptr, d_work, info_ptr);
315:       #endif
316:     PetscCheck(hip_status == rocblas_status_success, PETSC_COMM_SELF, PETSC_ERR_LIB, "rocsolver_*syevd failed for batch %" PetscInt_FMT, (PetscInt)i);
317:   }

319:   /* Check info */
320:   PetscCall(PetscMalloc1(n_batch, &h_info));
321:   PetscCallHIP(hipMemcpy(h_info, d_info, sizeof(int) * n_batch, hipMemcpyDeviceToHost));
322:   for (PetscInt i = 0; i < n_batch; i++) PetscCheck(h_info[i] == 0, PETSC_COMM_SELF, PETSC_ERR_LIB, "rocSOLVER eigendecomposition failed for matrix %" PetscInt_FMT ": info=%d", i, h_info[i]);
323:   PetscCall(PetscFree(h_info));

325:   /* Copy results back from contiguous layout to V_batch */
326:   Kokkos::parallel_for(
327:     "CopyResultsBack", Kokkos::RangePolicy<Kokkos::DefaultExecutionSpace>(0, n_batch), KOKKOS_LAMBDA(const int i) {
328:       for (int j = 0; j < n_size; j++) {
329:         for (int k = 0; k < n_size; k++) V_batch(i, j, k) = d_A_contig[i * n_size * n_size + k * n_size + j];
330:         Lambda_batch(i, j) = d_W_contig[i * n_size + j];
331:       }
332:     });
333:   Kokkos::fence();
334:     #endif
335:   PetscFunctionReturn(PETSC_SUCCESS);
336: }
337:   #elif defined(KOKKOS_ENABLE_SYCL)
338: static PetscErrorCode BatchedEigenSolve_Device(LETKFView3D T_batch, LETKFView2D Lambda_batch, LETKFView3D V_batch, PetscInt n_batch, PetscInt n_size, sycl::queue *q, EigenWorkspace *work)
339: {
340:   PetscFunctionBegin;
341:     #if PetscDefined(USE_COMPLEX)
342:   /* oneMKL's syevd USM overload targets real symmetric matrices; the complex analogue is heevd.
343:      The dispatcher gates the Kokkos path off when PETSC_USE_COMPLEX is set, so this is unreachable
344:      in practice; SETERRQ here as defense-in-depth in case that gate ever changes. */
345:   SETERRQ(PETSC_COMM_SELF, PETSC_ERR_SUP, "Complex numbers not supported on SYCL backend for LETKF");
346:     #else
347:   /* Use pre-allocated workspace */
348:   PetscScalar *d_work     = work->d_work;
349:   PetscScalar *d_A_contig = work->d_A_contig;
350:   PetscScalar *d_W_contig = work->d_W_contig;

352:   /* Copy T_batch to contiguous layout for oneMKL */
353:   Kokkos::parallel_for(
354:     "ReorganizeForOneMKL", Kokkos::RangePolicy<Kokkos::DefaultExecutionSpace>(0, n_batch), KOKKOS_LAMBDA(const int i) {
355:       for (int j = 0; j < n_size; j++) {
356:         for (int k = 0; k < n_size; k++) d_A_contig[i * n_size * n_size + k * n_size + j] = T_batch(i, j, k);
357:       }
358:     });
359:   Kokkos::fence();

361:   /* oneMKL doesn't have a native batched syevd, so we loop over batch and call the USM
362:      overload of oneapi::mkl::lapack::syevd. The USM overload reports failures via
363:      oneapi::mkl::lapack::lapack_exception (a sycl::exception subclass), not through an
364:      output `info` parameter; PetscCallCXX() catches std::exception and converts to a
365:      PETSc error. */
366:   for (int i = 0; i < n_batch; i++) {
367:     PetscScalar *A_ptr = d_A_contig + i * n_size * n_size;
368:     PetscScalar *W_ptr = d_W_contig + i * n_size;

370:     PetscCallCXX(oneapi::mkl::lapack::syevd(*q, oneapi::mkl::job::vec, oneapi::mkl::uplo::upper, n_size, A_ptr, n_size, W_ptr, d_work, work->lwork_device));
371:     PetscCallCXX(q->wait_and_throw());
372:   }

374:   /* Copy results back from contiguous layout to V_batch */
375:   Kokkos::parallel_for(
376:     "CopyResultsBack", Kokkos::RangePolicy<Kokkos::DefaultExecutionSpace>(0, n_batch), KOKKOS_LAMBDA(const int i) {
377:       for (int j = 0; j < n_size; j++) {
378:         for (int k = 0; k < n_size; k++) V_batch(i, j, k) = d_A_contig[i * n_size * n_size + k * n_size + j];
379:         Lambda_batch(i, j) = d_W_contig[i * n_size + j];
380:       }
381:     });
382:   Kokkos::fence();
383:     #endif
384:   PetscFunctionReturn(PETSC_SUCCESS);
385: }
386:   #endif
387: #endif

389: /*
390:   BatchedEigenSolve - Compute eigendecomposition for a batch of symmetric matrices

392:   Input Parameters:
393: + T_batch      - batch of symmetric matrices (n_batch x n_size x n_size)
394: . n_batch      - number of matrices in the batch
395: - n_size       - size of each matrix (m x m)
396: - device_handle - device-specific solver handle (only for device builds)
397: - work         - reusable workspace structure

399:   Output Parameters:
400: + Lambda_batch - eigenvalues for each matrix (n_batch x n_size)
401: - V_batch      - eigenvectors for each matrix (n_batch x n_size x n_size)

403:   Notes:
404:   Dispatcher function that calls the appropriate backend (Device or Host).
405: */
406: #if defined(KOKKOS_ENABLE_CUDA) || defined(KOKKOS_ENABLE_HIP) || defined(KOKKOS_ENABLE_SYCL)
407:   #if defined(KOKKOS_ENABLE_CUDA)
408: static PetscErrorCode BatchedEigenSolve(LETKFView3D T_batch, LETKFView2D Lambda_batch, LETKFView3D V_batch, PetscInt n_batch, PetscInt n_size, cusolverDnHandle_t device_handle, EigenWorkspace *work)
409: {
410:   PetscFunctionBegin;
411:   PetscCall(BatchedEigenSolve_Device(T_batch, Lambda_batch, V_batch, n_batch, n_size, device_handle, work));
412:   PetscFunctionReturn(PETSC_SUCCESS);
413: }
414:   #elif defined(KOKKOS_ENABLE_HIP)
415: static PetscErrorCode BatchedEigenSolve(LETKFView3D T_batch, LETKFView2D Lambda_batch, LETKFView3D V_batch, PetscInt n_batch, PetscInt n_size, rocblas_handle device_handle, EigenWorkspace *work)
416: {
417:   PetscFunctionBegin;
418:   PetscCall(BatchedEigenSolve_Device(T_batch, Lambda_batch, V_batch, n_batch, n_size, device_handle, work));
419:   PetscFunctionReturn(PETSC_SUCCESS);
420: }
421:   #elif defined(KOKKOS_ENABLE_SYCL)
422: static PetscErrorCode BatchedEigenSolve(LETKFView3D T_batch, LETKFView2D Lambda_batch, LETKFView3D V_batch, PetscInt n_batch, PetscInt n_size, sycl::queue *device_handle, EigenWorkspace *work)
423: {
424:   PetscFunctionBegin;
425:   PetscCall(BatchedEigenSolve_Device(T_batch, Lambda_batch, V_batch, n_batch, n_size, device_handle, work));
426:   PetscFunctionReturn(PETSC_SUCCESS);
427: }
428:   #endif
429: #else
430: static PetscErrorCode BatchedEigenSolve(LETKFView3D T_batch, LETKFView2D Lambda_batch, LETKFView3D V_batch, PetscInt n_batch, PetscInt n_size, EigenWorkspace *work)
431: {
432:   PetscFunctionBegin;
433:   PetscCall(BatchedEigenSolve_Host(T_batch, Lambda_batch, V_batch, n_batch, n_size, work));
434:   PetscFunctionReturn(PETSC_SUCCESS);
435: }
436: #endif

438: /*
439:   PetscDALETKFSetupLocalization_Kokkos - Prepares device views for localization matrix Q
440: */
441: PETSC_INTERN PetscErrorCode PetscDALETKFSetupLocalization_Kokkos(PetscDA_LETKF *impl)
442: {
443:   PetscInt nrows, rstart, rend, i, nnz, total_nnz;

445:   PetscFunctionBegin;
446:   PetscCheck(impl->Q, PETSC_COMM_SELF, PETSC_ERR_ARG_WRONGSTATE, "impl->Q is not set; PetscDALETKFInstallQ() must run before SetupLocalization");
447:   PetscCall(PetscKokkosInitializeCheck());

449:   PetscCall(MatGetOwnershipRange(impl->Q, &rstart, &rend));
450:   nrows = rend - rstart;
451:   /* impl->n_nnz_local was populated by PetscDALETKFInstallQ() from MatGetInfo(MAT_LOCAL) so we
452:      don't re-query (which would force a device->host sync on AIJKOKKOS). */
453:   total_nnz = impl->n_nnz_local;

455:   /* Define View types */
456:   using view_1d_int    = Kokkos::View<PetscInt *, Kokkos::LayoutLeft>;
457:   using view_1d_scalar = Kokkos::View<PetscScalar *, Kokkos::LayoutLeft>;

459:   /* Allocate device views using actual total nnz from Q */
460:   view_1d_int    *d_Q_i;
461:   view_1d_int    *d_Q_j;
462:   view_1d_scalar *d_Q_a;

464:   PetscCallCXX(d_Q_i = new view_1d_int("Q_i", nrows + 1));
465:   PetscCallCXX(d_Q_j = new view_1d_int("Q_j", total_nnz));
466:   PetscCallCXX(d_Q_a = new view_1d_scalar("Q_a", total_nnz));

468:   /* Create host mirrors */
469:   Kokkos::View<PetscInt *, Kokkos::LayoutLeft, Kokkos::HostSpace>    h_Q_i;
470:   Kokkos::View<PetscInt *, Kokkos::LayoutLeft, Kokkos::HostSpace>    h_Q_j;
471:   Kokkos::View<PetscScalar *, Kokkos::LayoutLeft, Kokkos::HostSpace> h_Q_a;
472:   PetscCallCXX(h_Q_i = Kokkos::create_mirror_view(Kokkos::HostSpace(), *d_Q_i));
473:   PetscCallCXX(h_Q_j = Kokkos::create_mirror_view(Kokkos::HostSpace(), *d_Q_j));
474:   PetscCallCXX(h_Q_a = Kokkos::create_mirror_view(Kokkos::HostSpace(), *d_Q_a));

476:   /* Fill host mirrors with LOCAL indices into obs_work */
477:   h_Q_i(0) = 0;
478:   for (i = 0; i < nrows; i++) {
479:     const PetscInt    *cols;
480:     const PetscScalar *vals;
481:     PetscCall(MatGetRow(impl->Q, rstart + i, &nnz, &cols, &vals));
482:     h_Q_i(i + 1) = h_Q_i(i) + nnz;
483:     for (PetscInt k = 0; k < nnz; k++) {
484:       PetscInt local_idx;
485:       PetscCall(PetscHMapIGet(impl->obs_g2l, cols[k], &local_idx));
486:       PetscCheck(local_idx >= 0, PETSC_COMM_SELF, PETSC_ERR_PLIB, "Observation index %" PetscInt_FMT " not found in local map", cols[k]);
487:       h_Q_j(h_Q_i(i) + k) = local_idx;
488:       h_Q_a(h_Q_i(i) + k) = vals[k];
489:     }
490:     PetscCall(MatRestoreRow(impl->Q, rstart + i, &nnz, &cols, &vals));
491:   }

493:   /* Copy to device */
494:   PetscCallCXX(Kokkos::deep_copy(*d_Q_i, h_Q_i));
495:   PetscCallCXX(Kokkos::deep_copy(*d_Q_j, h_Q_j));
496:   PetscCallCXX(Kokkos::deep_copy(*d_Q_a, h_Q_a));

498:   /* Store in impl */
499:   PetscCheck(!impl->Q_device_i, PETSC_COMM_SELF, PETSC_ERR_ARG_WRONGSTATE, "Q_device_i already allocated; PetscDALETKFDestroyLocalization_Kokkos must run before re-setup");
500:   impl->Q_device_i = static_cast<void *>(d_Q_i);
501:   impl->Q_device_j = static_cast<void *>(d_Q_j);
502:   impl->Q_device_a = static_cast<void *>(d_Q_a);
503:   PetscFunctionReturn(PETSC_SUCCESS);
504: }

506: /*
507:   PetscDALETKFDestroyQDeviceMirrors_Kokkos - Free only the device-side CSR mirrors of Q.

509:   Used on the Q-rebuild path (setters that mutate type/radius/coordinates) so the persistent
510:   eigensolver workspace and the cusolver/rocblas/SYCL handle survive across rebuilds. The
511:   full destroy below also calls this helper.
512: */
513: PETSC_INTERN PetscErrorCode PetscDALETKFDestroyQDeviceMirrors_Kokkos(PetscDA_LETKF *impl)
514: {
515:   PetscFunctionBegin;
516:   if (impl->Q_device_i) {
517:     using view_1d_int = Kokkos::View<PetscInt *, Kokkos::LayoutLeft>;
518:     delete static_cast<view_1d_int *>(impl->Q_device_i);
519:     impl->Q_device_i = NULL;
520:   }
521:   if (impl->Q_device_j) {
522:     using view_1d_int = Kokkos::View<PetscInt *, Kokkos::LayoutLeft>;
523:     delete static_cast<view_1d_int *>(impl->Q_device_j);
524:     impl->Q_device_j = NULL;
525:   }
526:   if (impl->Q_device_a) {
527:     using view_1d_scalar = Kokkos::View<PetscScalar *, Kokkos::LayoutLeft>;
528:     delete static_cast<view_1d_scalar *>(impl->Q_device_a);
529:     impl->Q_device_a = NULL;
530:   }
531:   PetscFunctionReturn(PETSC_SUCCESS);
532: }

534: /*
535:   PetscDALETKFDestroyLocalization_Kokkos - Free all device-side state owned by the Kokkos backend.

537:   Tears down the Q device mirrors AND the persistent eigensolver workspace + cusolver/rocblas/SYCL
538:   handle. Both LOC_NONE (GlobalAnalysis_Kokkos) and the per-vertex paths allocate the latter
539:   state, so PetscDADestroy_LETKF calls this regardless of localization type. Q-rebuild paths
540:   use PetscDALETKFDestroyQDeviceMirrors_Kokkos() instead so the handle and workspace persist.
541: */
542: PETSC_INTERN PetscErrorCode PetscDALETKFDestroyLocalization_Kokkos(PetscDA_LETKF *impl)
543: {
544:   PetscFunctionBegin;
545:   PetscCall(PetscDALETKFDestroyQDeviceMirrors_Kokkos(impl));

547:   /* Destroy solver handle and workspace */
548:   if (impl->eigen_work) {
549:     EigenWorkspace *work = static_cast<EigenWorkspace *>(impl->eigen_work);

551: #if defined(KOKKOS_ENABLE_CUDA) || defined(KOKKOS_ENABLE_HIP) || defined(KOKKOS_ENABLE_SYCL)
552:   #if defined(KOKKOS_ENABLE_CUDA)
553:     PetscCallCUDA(cudaFree(work->d_A_contig));
554:     PetscCallCUDA(cudaFree(work->d_W_contig));
555:     PetscCallCUDA(cudaFree(work->d_work));
556:     PetscCallCUDA(cudaFree(work->d_info));
557:     /* Destroy returns ignored: teardown may race with Kokkos/CUDA context shutdown when the
558:        enclosing PetscDA outlives PetscFinalize() handlers; raising here would mask the real
559:        teardown order issue. */
560:     if (work->syevj_params) cusolverDnDestroySyevjInfo(work->syevj_params);
561:   #elif defined(KOKKOS_ENABLE_HIP)
562:     PetscCallHIP(hipFree(work->d_A_contig));
563:     PetscCallHIP(hipFree(work->d_W_contig));
564:     PetscCallHIP(hipFree(work->d_work));
565:     PetscCallHIP(hipFree(work->d_info));
566:   #elif defined(KOKKOS_ENABLE_SYCL)
567:     if (impl->solver_handle) {
568:       sycl::queue *q = static_cast<sycl::queue *>(impl->solver_handle);
569:       if (work->d_A_contig) sycl::free(work->d_A_contig, *q);
570:       if (work->d_W_contig) sycl::free(work->d_W_contig, *q);
571:       if (work->d_work) sycl::free(work->d_work, *q);
572:       if (work->d_info) sycl::free(work->d_info, *q);
573:     }
574:   #endif
575: #else
576:   #if PetscDefined(USE_COMPLEX)
577:     PetscCall(PetscFree4(work->all_v, work->all_lambda, work->all_work, work->all_rwork));
578:   #else
579:     PetscCall(PetscFree3(work->all_v, work->all_lambda, work->all_work));
580:   #endif
581: #endif

583:     delete work;
584:     impl->eigen_work = NULL;
585:   }

587:   if (impl->solver_handle) {
588:     /* Destroy returns ignored: see comment above on teardown-time race with Kokkos/CUDA
589:        context shutdown. */
590: #if defined(KOKKOS_ENABLE_CUDA)
591:     cusolverDnDestroy(static_cast<cusolverDnHandle_t>(impl->solver_handle));
592: #elif defined(KOKKOS_ENABLE_HIP)
593:     rocblas_destroy_handle(static_cast<rocblas_handle>(impl->solver_handle));
594: #elif defined(KOKKOS_ENABLE_SYCL)
595:     delete static_cast<sycl::queue *>(impl->solver_handle);
596: #endif
597:     impl->solver_handle = NULL;
598:   }
599:   PetscFunctionReturn(PETSC_SUCCESS);
600: }

602: /* ========================================================================== */
603: /*                    LETKF Local Analysis (Main Function)                    */
604: /* ========================================================================== */

606: /*
607:   PetscDALETKFLocalAnalysis_Kokkos - Performs local LETKF analysis for all grid points (Kokkos version)

609:   Input Parameters:
610: + da             - the PetscDA context
611: . impl           - LETKF implementation data
612: . m              - ensemble size
613: . n_vertices     - number of grid points
614: . X              - global anomaly matrix (state_size x m)
615: . observation    - observation vector
616: . Z_global       - global observation ensemble (obs_size x m)
617: . y_mean_global  - global observation mean
618: - r_inv_sqrt_global - global R^{-1/2}

620:   Output:
621: . da->ensemble - updated with analysis ensemble

623:   Notes:
624:   Kokkos device implementation of the LETKF local analysis. All n_vertices grid points are
625:   processed in a batched fashion: per-vertex S, T, V, T_sqrt, and weight slabs live in
626:   device 3-D/2-D views, observation data is mirrored from host when needed, and the
627:   per-vertex eigendecompositions are dispatched through BatchedEigenSolve_Device() (or
628:   BatchedEigenSolve_Host() when Kokkos's default execution space is the host).
629: */
630: PETSC_INTERN PetscErrorCode PetscDALETKFLocalAnalysis_Kokkos(PetscDA da, PetscDA_LETKF *impl, PetscInt m, PetscInt n_vertices, Mat X, Vec observation, Mat Z_global, Vec y_mean_global, Vec r_inv_sqrt_global)
631: {
632:   using exec_space              = Kokkos::DefaultExecutionSpace;
633:   using view_3d                 = Kokkos::View<PetscScalar ***, Kokkos::LayoutLeft, exec_space>;
634:   using view_2d                 = Kokkos::View<PetscScalar **, Kokkos::LayoutLeft, exec_space>;
635:   using view_1d_int_const       = Kokkos::View<const PetscInt *, Kokkos::LayoutLeft, Kokkos::MemoryTraits<Kokkos::Unmanaged>>;
636:   using view_1d_scalar_const    = Kokkos::View<const PetscScalar *, Kokkos::LayoutLeft, Kokkos::MemoryTraits<Kokkos::Unmanaged>>;
637:   using view_1d_int             = Kokkos::View<PetscInt *, Kokkos::LayoutLeft>;
638:   using view_1d_scalar          = Kokkos::View<PetscScalar *, Kokkos::LayoutLeft>;
639:   using view_2d_unmanaged       = Kokkos::View<const PetscScalar **, Kokkos::LayoutLeft, Kokkos::MemoryTraits<Kokkos::Unmanaged>>;
640:   using view_1d_unmanaged       = Kokkos::View<const PetscScalar *, Kokkos::LayoutLeft, Kokkos::MemoryTraits<Kokkos::Unmanaged>>;
641:   using view_2d_unmanaged_write = Kokkos::View<PetscScalar **, Kokkos::LayoutLeft, Kokkos::MemoryTraits<Kokkos::Unmanaged>>;
642:   PetscDA_Ensemble       *en    = &impl->en;
643:   EigenWorkspace         *eigen_work;
644:   PetscInt                ndof;
645:   PetscInt                lda_z_global, lda_x, lda_e, n_obs_local;
646:   PetscInt                max_nnz_per_row, max_nnz_copy, chunk_size;
647:   PetscInt64              mem_per_point;
648:   PetscReal               sqrt_m_minus_1, scale, inflation_inv;
649:   PetscReal               flops, n_obs_total;
650:   PetscMemType            z_mem_type, y_mem_type, y_mean_mem_type, r_inv_sqrt_mem_type;
651:   PetscMemType            x_mem_type, mean_mem_type, e_mem_type;
652:   const PetscScalar      *z_global_array, *y_global_array, *y_mean_global_array, *r_inv_sqrt_global_array;
653:   const PetscScalar      *x_array, *mean_array;
654:   PetscScalar            *e_array;
655:   const PetscScalar      *z_ptr, *y_ptr, *y_mean_ptr, *r_inv_sqrt_ptr;
656:   const PetscScalar      *x_ptr, *mean_ptr;
657:   PetscScalar            *e_ptr;
658:   PetscBool               e_is_copy = PETSC_FALSE;
659:   view_1d_int_const       Q_i_view, Q_j_view;
660:   view_1d_scalar_const    Q_a_view;
661:   LETKFView2D             z_managed, x_managed, e_managed;
662:   LETKFView1D             y_managed, y_mean_managed, r_inv_sqrt_managed, mean_managed;
663:   view_2d_unmanaged       Z_global_view, X_view;
664:   view_1d_unmanaged       y_global_view, y_mean_global_view, r_inv_sqrt_global_view, mean_view;
665:   view_2d_unmanaged_write E_view;
666:   view_3d                 S_batch, T_batch, V_batch, T_sqrt_batch;
667:   view_2d                 Lambda_batch, w_batch, delta_batch, y_batch, y_mean_batch, r_inv_sqrt_batch, temp1_batch, temp2_batch, inv_sqrt_lambda_batch;
668: #if defined(KOKKOS_ENABLE_CUDA)
669:   cusolverDnHandle_t device_handle = nullptr;
670:   cusolverStatus_t   cusolver_status;
671: #elif defined(KOKKOS_ENABLE_HIP)
672:   rocblas_handle device_handle = nullptr;
673: #elif defined(KOKKOS_ENABLE_SYCL)
674:   sycl::queue *device_handle = nullptr;
675: #endif

677:   PetscFunctionBegin;
678:   ndof           = da->ndof;
679:   scale          = 1.0 / PetscSqrtReal((PetscReal)(m - 1));
680:   sqrt_m_minus_1 = PetscSqrtReal((PetscReal)(m - 1));
681:   inflation_inv  = 1.0 / en->inflation; /* (1/rho) for T matrix: T = (1/rho)I + S^T*S */

683:   /* ===================================================================== */
684:   /* Step 2.1.1: Create batched workspace for ALL grid points            */
685:   /* ===================================================================== */
686:   /*
687:      NOTE ON PARALLELISM STRATEGY:
688:      We use Kokkos::RangePolicy over grid points (n_vertices) combined with KokkosBatched::Serial kernels.
689:      Since the data layout is LayoutLeft (Column-Major) to match PETSc/LAPACK, the index 'i' (grid point)
690:      is the fastest varying index (stride 1).

692:      RangePolicy maps consecutive threads to consecutive 'i', ensuring perfect memory coalescing
693:      when accessing arrays like S_batch(i, p, j).

695:      Using TeamPolicy/TeamVectorRange to parallelize inner loops (m or p) would assign a team to 'i',
696:      causing threads within the team to access S_batch with stride 'n_vertices', which leads to
697:      uncoalesced memory access and poor performance on GPUs.

699:      Therefore, RangePolicy + SerialGemm is the optimal strategy for this data layout.
700:   */

702:   /* ===================================================================== */
703:   /* Step 2.1.2a: Pre-extract Q matrix CSR data for device access        */
704:   /* ===================================================================== */
705:   PetscCheck(impl->Q_device_i, PETSC_COMM_SELF, PETSC_ERR_ARG_WRONGSTATE, "Q device views not allocated; PetscDALETKFSetupLocalization_Kokkos must run before LocalAnalysis");
706:   /* Use pre-allocated device views */
707:   view_1d_int    *d_Q_i = static_cast<view_1d_int *>(impl->Q_device_i);
708:   view_1d_int    *d_Q_j = static_cast<view_1d_int *>(impl->Q_device_j);
709:   view_1d_scalar *d_Q_a = static_cast<view_1d_scalar *>(impl->Q_device_a);

711:   Q_i_view = view_1d_int_const(d_Q_i->data(), d_Q_i->extent(0));
712:   Q_j_view = view_1d_int_const(d_Q_j->data(), d_Q_j->extent(0));
713:   Q_a_view = view_1d_scalar_const(d_Q_a->data(), d_Q_a->extent(0));

715:   /* Get global observation data arrays */
716:   PetscCall(MatDenseGetArrayReadAndMemType(Z_global, &z_global_array, &z_mem_type));
717:   PetscCall(VecGetArrayReadAndMemType(observation, &y_global_array, &y_mem_type));
718:   PetscCall(VecGetArrayReadAndMemType(y_mean_global, &y_mean_global_array, &y_mean_mem_type));
719:   PetscCall(VecGetArrayReadAndMemType(r_inv_sqrt_global, &r_inv_sqrt_global_array, &r_inv_sqrt_mem_type));
720:   PetscCall(MatDenseGetLDA(Z_global, &lda_z_global));
721:   PetscCall(VecGetLocalSize(observation, &n_obs_local));

723:   /* Handle memory mirroring for observation data. The 1-D obs vectors are sized n_obs_local;
724:      only the 2-D Z view is shaped by lda_z_global. */
725:   z_ptr          = z_global_array;
726:   y_ptr          = y_global_array;
727:   y_mean_ptr     = y_mean_global_array;
728:   r_inv_sqrt_ptr = r_inv_sqrt_global_array;

730:   if (z_mem_type == PETSC_MEMTYPE_HOST) {
731:     Kokkos::View<const PetscScalar **, Kokkos::LayoutLeft, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>> src(z_global_array, lda_z_global, m);
732:     PetscCallCXX(z_managed = LETKFView2D("z_managed", lda_z_global, m));
733:     PetscCallCXX(Kokkos::deep_copy(z_managed, src));
734:     z_ptr = z_managed.data();
735:   }
736:   if (y_mem_type == PETSC_MEMTYPE_HOST) {
737:     Kokkos::View<const PetscScalar *, Kokkos::LayoutLeft, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>> src(y_global_array, n_obs_local);
738:     PetscCallCXX(y_managed = LETKFView1D("y_managed", n_obs_local));
739:     PetscCallCXX(Kokkos::deep_copy(y_managed, src));
740:     y_ptr = y_managed.data();
741:   }
742:   if (y_mean_mem_type == PETSC_MEMTYPE_HOST) {
743:     Kokkos::View<const PetscScalar *, Kokkos::LayoutLeft, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>> src(y_mean_global_array, n_obs_local);
744:     PetscCallCXX(y_mean_managed = LETKFView1D("y_mean_managed", n_obs_local));
745:     PetscCallCXX(Kokkos::deep_copy(y_mean_managed, src));
746:     y_mean_ptr = y_mean_managed.data();
747:   }
748:   if (r_inv_sqrt_mem_type == PETSC_MEMTYPE_HOST) {
749:     Kokkos::View<const PetscScalar *, Kokkos::LayoutLeft, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>> src(r_inv_sqrt_global_array, n_obs_local);
750:     PetscCallCXX(r_inv_sqrt_managed = LETKFView1D("r_inv_sqrt_managed", n_obs_local));
751:     PetscCallCXX(Kokkos::deep_copy(r_inv_sqrt_managed, src));
752:     r_inv_sqrt_ptr = r_inv_sqrt_managed.data();
753:   }

755:   /* Create unmanaged Kokkos views for global observation data. NOTE: Z_global_view's leading
756:      extent is lda_z_global (the MatDense column stride) while the 1-D obs views use n_obs_local
757:      (the unpadded local-row count). Always index Z_global_view as Z_global_view(obs_idx, j) with
758:      obs_idx < n_obs_local taken from Q's per-row column list; never iterate [0, extent(0)) on Z
759:      because that walks into LDA padding. */
760:   Z_global_view          = view_2d_unmanaged(z_ptr, lda_z_global, m);
761:   y_global_view          = view_1d_unmanaged(y_ptr, n_obs_local);
762:   y_mean_global_view     = view_1d_unmanaged(y_mean_ptr, n_obs_local);
763:   r_inv_sqrt_global_view = view_1d_unmanaged(r_inv_sqrt_ptr, n_obs_local);

765:   /* Get access to global X matrix and mean vector */
766:   PetscCall(MatDenseGetArrayReadAndMemType(X, &x_array, &x_mem_type));
767:   PetscCall(VecGetArrayReadAndMemType(impl->mean, &mean_array, &mean_mem_type));
768:   PetscCall(MatDenseGetArrayWriteAndMemType(en->ensemble, &e_array, &e_mem_type));
769:   PetscCall(MatDenseGetLDA(X, &lda_x));
770:   PetscCall(MatDenseGetLDA(en->ensemble, &lda_e));
771:   /* Per-vertex subviews into X/E span [i_global*ndof, (i_global+1)*ndof); the maximum global
772:      index is n_vertices-1, so the leading dimension must cover all local vertices. */
773:   PetscCheck(lda_x >= n_vertices * ndof, PetscObjectComm((PetscObject)X), PETSC_ERR_ARG_INCOMP, "X leading dimension %" PetscInt_FMT " < n_vertices*ndof %" PetscInt_FMT, lda_x, n_vertices * ndof);
774:   PetscCheck(lda_e >= n_vertices * ndof, PetscObjectComm((PetscObject)en->ensemble), PETSC_ERR_ARG_INCOMP, "Ensemble leading dimension %" PetscInt_FMT " < n_vertices*ndof %" PetscInt_FMT, lda_e, n_vertices * ndof);

776:   /* Handle memory mirroring for state data */
777:   x_ptr    = x_array;
778:   mean_ptr = mean_array;
779:   e_ptr    = e_array;

781:   if (x_mem_type == PETSC_MEMTYPE_HOST) {
782:     Kokkos::View<const PetscScalar **, Kokkos::LayoutLeft, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>> src(x_array, lda_x, m);
783:     PetscCallCXX(x_managed = LETKFView2D("x_managed", lda_x, m));
784:     PetscCallCXX(Kokkos::deep_copy(x_managed, src));
785:     x_ptr = x_managed.data();
786:   }
787:   if (mean_mem_type == PETSC_MEMTYPE_HOST) {
788:     /* impl->mean is a Vec with local size n_vertices*ndof (no MatDense LDA padding), so size the
789:        mirror to the exact buffer extent; reading lda_x would over-read when MatDense pads X's LDA. */
790:     Kokkos::View<const PetscScalar *, Kokkos::LayoutLeft, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>> src(mean_array, n_vertices * ndof);
791:     PetscCallCXX(mean_managed = LETKFView1D("mean_managed", n_vertices * ndof));
792:     PetscCallCXX(Kokkos::deep_copy(mean_managed, src));
793:     mean_ptr = mean_managed.data();
794:   }
795:   if (e_mem_type == PETSC_MEMTYPE_HOST) {
796:     Kokkos::View<PetscScalar **, Kokkos::LayoutLeft, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>> src(e_array, lda_e, m);
797:     PetscCallCXX(e_managed = LETKFView2D("e_managed", lda_e, m));
798:     PetscCallCXX(Kokkos::deep_copy(e_managed, src));
799:     e_ptr     = e_managed.data();
800:     e_is_copy = PETSC_TRUE;
801:   }

803:   /* Create unmanaged Kokkos views for global data */
804:   X_view    = view_2d_unmanaged(const_cast<PetscScalar *>(x_ptr), lda_x, m);
805:   mean_view = view_1d_unmanaged(mean_ptr, n_vertices * ndof);
806:   E_view    = view_2d_unmanaged_write(e_ptr, lda_e, m);

808:   max_nnz_per_row = impl->max_nnz_per_row;

810:   /* Determine chunk size to avoid OOM on large grids */
811:   mem_per_point = (PetscInt64)sizeof(PetscScalar) * ((PetscInt64)m * m + (PetscInt64)max_nnz_per_row * m);
812:   if (impl->batch_size > 0) {
813:     chunk_size = impl->batch_size;
814:   } else {
815:     /* Target ~2GB workspace. Approx memory per point: m*m*sizeof(PetscScalar) (T) + p*m*sizeof(PetscScalar) (Z). */
816:     chunk_size = (PetscInt)((PetscInt64)2 * 1024 * 1024 * 1024 / mem_per_point);
817:     /* Clamp to reasonable max to avoid huge allocations even if memory allows */
818:     if (chunk_size > 32768) chunk_size = 32768;
819:   }

821:   if (chunk_size < 1) chunk_size = 1;
822:   if (chunk_size > n_vertices) chunk_size = n_vertices;

824:   /* OPTIMIZATION: Create device solver handle once, reuse across chunks */
825: #if defined(KOKKOS_ENABLE_CUDA)
826:   if (impl->solver_handle) device_handle = static_cast<cusolverDnHandle_t>(impl->solver_handle);
827:   else {
828:     cusolver_status = cusolverDnCreate(&device_handle);
829:     PetscCheck(cusolver_status == CUSOLVER_STATUS_SUCCESS, PETSC_COMM_SELF, PETSC_ERR_LIB, "cusolverDnCreate failed");
830:     impl->solver_handle = static_cast<void *>(device_handle);
831:   }
832: #elif defined(KOKKOS_ENABLE_HIP)
833:   if (impl->solver_handle) device_handle = static_cast<rocblas_handle>(impl->solver_handle);
834:   else {
835:     rocblas_status hip_status = rocblas_create_handle(&device_handle);
836:     PetscCheck(hip_status == rocblas_status_success, PETSC_COMM_SELF, PETSC_ERR_LIB, "rocblas_create_handle failed");
837:     impl->solver_handle = static_cast<void *>(device_handle);
838:   }
839: #elif defined(KOKKOS_ENABLE_SYCL)
840:   if (impl->solver_handle) device_handle = static_cast<sycl::queue *>(impl->solver_handle);
841:   else {
842:     PetscCallCXX(device_handle = new sycl::queue(sycl::gpu_selector_v));
843:     impl->solver_handle = static_cast<void *>(device_handle);
844:   }
845: #endif

847:   /* ===================================================================== */
848:   /* OPTIMIZATION: Hoist allocations outside the chunk loop                */
849:   /* ===================================================================== */
850:   /* Allocate Kokkos Views once for the maximum chunk size */
851:   max_nnz_copy = max_nnz_per_row;

853:   eigen_work = static_cast<EigenWorkspace *>(impl->eigen_work);
854:   if (!eigen_work) {
855:     PetscCallCXX(eigen_work = new EigenWorkspace());
856:     impl->eigen_work = static_cast<void *>(eigen_work);
857:   }

859:   /* Check if reallocation is needed */
860:   if (eigen_work->max_chunk_size < chunk_size || eigen_work->m != m || eigen_work->max_nnz < max_nnz_copy) {
861:     /* Free old device workspace if exists */
862: #if defined(KOKKOS_ENABLE_CUDA)
863:     PetscCallCUDA(cudaFree(eigen_work->d_work));
864:     PetscCallCUDA(cudaFree(eigen_work->d_info));
865:     PetscCallCUDA(cudaFree(eigen_work->d_A_contig));
866:     PetscCallCUDA(cudaFree(eigen_work->d_W_contig));
867:     if (eigen_work->syevj_params) cusolverDnDestroySyevjInfo(eigen_work->syevj_params);
868:     eigen_work->syevj_params = nullptr;
869: #elif defined(KOKKOS_ENABLE_HIP)
870:     PetscCallHIP(hipFree(eigen_work->d_work));
871:     PetscCallHIP(hipFree(eigen_work->d_info));
872:     PetscCallHIP(hipFree(eigen_work->d_A_contig));
873:     PetscCallHIP(hipFree(eigen_work->d_W_contig));
874: #elif defined(KOKKOS_ENABLE_SYCL)
875:     if (eigen_work->d_work) sycl::free(eigen_work->d_work, *device_handle);
876:     if (eigen_work->d_info) sycl::free(eigen_work->d_info, *device_handle);
877:     if (eigen_work->d_A_contig) sycl::free(eigen_work->d_A_contig, *device_handle);
878:     if (eigen_work->d_W_contig) sycl::free(eigen_work->d_W_contig, *device_handle);
879: #endif

881: #if !defined(KOKKOS_ENABLE_CUDA) && !defined(KOKKOS_ENABLE_HIP) && !defined(KOKKOS_ENABLE_SYCL)
882:   #if PetscDefined(USE_COMPLEX)
883:     if (eigen_work->all_v) PetscCall(PetscFree4(eigen_work->all_v, eigen_work->all_lambda, eigen_work->all_work, eigen_work->all_rwork));
884:   #else
885:     if (eigen_work->all_v) PetscCall(PetscFree3(eigen_work->all_v, eigen_work->all_lambda, eigen_work->all_work));
886:   #endif
887: #endif

889:     /* Update dimensions */
890:     eigen_work->max_chunk_size = chunk_size;
891:     eigen_work->m              = m;
892:     eigen_work->max_nnz        = max_nnz_copy;

894:     /* Allocate Kokkos Views */
895:     PetscCallCXX(eigen_work->S_batch = view_3d("S_batch", chunk_size, max_nnz_copy, m));
896:     PetscCallCXX(eigen_work->T_batch = view_3d("T_batch", chunk_size, m, m));
897:     /* Alias: the eigensolve overwrites T in place, so V and T share storage. Any future
898:        kernel that needs the original symmetric T after the eigensolve must allocate V
899:        separately (view_3d("V_batch", chunk_size, m, m)) instead of aliasing. */
900:     eigen_work->V_batch = eigen_work->T_batch;
901:     PetscCallCXX(eigen_work->Lambda_batch = view_2d("Lambda_batch", chunk_size, m));
902:     PetscCallCXX(eigen_work->T_sqrt_batch = view_3d("T_sqrt_batch", chunk_size, m, m));
903:     PetscCallCXX(eigen_work->w_batch = view_2d("w_batch", chunk_size, m));
904:     PetscCallCXX(eigen_work->delta_batch = view_2d("delta_batch", chunk_size, max_nnz_copy));
905:     PetscCallCXX(eigen_work->y_batch = view_2d("y_batch", chunk_size, max_nnz_copy));
906:     PetscCallCXX(eigen_work->y_mean_batch = view_2d("y_mean_batch", chunk_size, max_nnz_copy));
907:     PetscCallCXX(eigen_work->r_inv_sqrt_batch = view_2d("r_inv_sqrt_batch", chunk_size, max_nnz_copy));
908:     PetscCallCXX(eigen_work->temp1_batch = view_2d("temp1_batch", chunk_size, m));
909:     PetscCallCXX(eigen_work->temp2_batch = view_2d("temp2_batch", chunk_size, m));
910:     PetscCallCXX(eigen_work->inv_sqrt_lambda_batch = view_2d("inv_sqrt_lambda_batch", chunk_size, m));

912:     /* Allocate solver workspace */
913: #if defined(KOKKOS_ENABLE_CUDA) || defined(KOKKOS_ENABLE_HIP) || defined(KOKKOS_ENABLE_SYCL)
914:   #if defined(KOKKOS_ENABLE_CUDA)
915:     {
916:       /* Create syevj params */
917:       cusolver_status = cusolverDnCreateSyevjInfo(&eigen_work->syevj_params);
918:       PetscCheck(cusolver_status == CUSOLVER_STATUS_SUCCESS, PETSC_COMM_SELF, PETSC_ERR_LIB, "cusolverDnCreateSyevjInfo failed");

920:       /* Set default params */
921:       cusolver_status = cusolverDnXsyevjSetTolerance(eigen_work->syevj_params, 1e-7);
922:       PetscCheck(cusolver_status == CUSOLVER_STATUS_SUCCESS, PETSC_COMM_SELF, PETSC_ERR_LIB, "cusolverDnXsyevjSetTolerance failed");
923:       cusolver_status = cusolverDnXsyevjSetMaxSweeps(eigen_work->syevj_params, 100);
924:       PetscCheck(cusolver_status == CUSOLVER_STATUS_SUCCESS, PETSC_COMM_SELF, PETSC_ERR_LIB, "cusolverDnXsyevjSetMaxSweeps failed");
925:       cusolver_status = cusolverDnXsyevjSetSortEig(eigen_work->syevj_params, 1); /* Sort eigenvalues */
926:       PetscCheck(cusolver_status == CUSOLVER_STATUS_SUCCESS, PETSC_COMM_SELF, PETSC_ERR_LIB, "cusolverDnXsyevjSetSortEig failed");

928:       /* Query workspace size */
929:       PetscScalar *d_A = eigen_work->T_batch.data();
930:       PetscScalar *d_W = eigen_work->Lambda_batch.data();
931:       int          lwork;
932:     #if PetscDefined(USE_REAL_SINGLE)
933:       cusolver_status = cusolverDnSsyevjBatched_bufferSize(device_handle, CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_UPPER, m, d_A, m, d_W, &lwork, eigen_work->syevj_params, chunk_size);
934:     #else
935:       cusolver_status = cusolverDnDsyevjBatched_bufferSize(device_handle, CUSOLVER_EIG_MODE_VECTOR, CUBLAS_FILL_MODE_UPPER, m, d_A, m, d_W, &lwork, eigen_work->syevj_params, chunk_size);
936:     #endif
937:       PetscCheck(cusolver_status == CUSOLVER_STATUS_SUCCESS, PETSC_COMM_SELF, PETSC_ERR_LIB, "cusolverDn*syevjBatched_bufferSize failed");
938:       eigen_work->lwork_device = lwork;

940:       /* Allocate workspace */
941:       PetscCallCUDA(cudaMalloc(&eigen_work->d_work, sizeof(PetscScalar) * lwork));
942:       PetscCallCUDA(cudaMalloc(&eigen_work->d_info, sizeof(int) * chunk_size));
943:       PetscCallCUDA(cudaMalloc(&eigen_work->d_A_contig, sizeof(PetscScalar) * chunk_size * m * m));
944:       PetscCallCUDA(cudaMalloc(&eigen_work->d_W_contig, sizeof(PetscScalar) * chunk_size * m));
945:     }
946:   #elif defined(KOKKOS_ENABLE_HIP)
947:     {
948:         /* rocsolver_*syevd takes a single n-element off-diagonal scratch buffer (the E array).
949:          The batch loop is sequential, so one shared buffer is sufficient. */
950:     #if PetscDefined(USE_COMPLEX)
951:       int lwork = 0; /* Complex not supported on device */
952:     #else
953:       PetscCheck(m <= INT_MAX, PETSC_COMM_SELF, PETSC_ERR_ARG_OUTOFRANGE, "Ensemble size m=%" PetscInt_FMT " exceeds INT_MAX for rocsolver lwork", m);
954:       int lwork = (int)m;
955:     #endif
956:       eigen_work->lwork_device = lwork;

958:       /* Allocate workspace */
959:       if (lwork > 0) {
960:         PetscCallHIP(hipMalloc(&eigen_work->d_work, sizeof(PetscScalar) * lwork));
961:         PetscCallHIP(hipMalloc(&eigen_work->d_info, sizeof(int) * chunk_size));
962:         PetscCallHIP(hipMalloc(&eigen_work->d_A_contig, sizeof(PetscScalar) * chunk_size * m * m));
963:         PetscCallHIP(hipMalloc(&eigen_work->d_W_contig, sizeof(PetscScalar) * chunk_size * m));
964:       }
965:     }
966:   #elif defined(KOKKOS_ENABLE_SYCL)
967:     {
968:       /* Query the exact scratchpad size oneMKL needs for syevd. The hand-rolled formula that
969:          used to live here was guessed from textbook LAPACK requirements and is not guaranteed
970:          to match every oneMKL backend. */
971:       std::int64_t lwork = 0;
972:       PetscCallCXX(lwork = oneapi::mkl::lapack::syevd_scratchpad_size<PetscScalar>(*device_handle, oneapi::mkl::job::vec, oneapi::mkl::uplo::upper, m, m));
973:       PetscCheck(lwork <= (std::int64_t)INT_MAX, PETSC_COMM_SELF, PETSC_ERR_PLIB, "oneMKL syevd_scratchpad_size %lld exceeds INT_MAX", (long long)lwork);
974:       eigen_work->lwork_device = (int)lwork;

976:       /* Allocate workspace using SYCL malloc_device. The USM overload of syevd reports failures
977:          via lapack_exception, not an info argument, so d_info is not allocated. */
978:       eigen_work->d_work     = sycl::malloc_device<PetscScalar>(lwork, *device_handle);
979:       eigen_work->d_A_contig = sycl::malloc_device<PetscScalar>(chunk_size * m * m, *device_handle);
980:       eigen_work->d_W_contig = sycl::malloc_device<PetscScalar>(chunk_size * m, *device_handle);
981:       eigen_work->d_info     = nullptr;
982:       PetscCheck(eigen_work->d_work && eigen_work->d_A_contig && eigen_work->d_W_contig, PETSC_COMM_SELF, PETSC_ERR_MEM, "SYCL memory allocation failed");
983:     }
984:   #endif
985: #else
986:     {
987:       PetscBLASInt n_blas;
988:       PetscCall(PetscBLASIntCast(m, &n_blas));
989:       eigen_work->n_blas = n_blas;

991:       /* Query workspace size */
992:       PetscBLASInt lwork_query = -1;
993:       PetscScalar  work_query;
994:       PetscBLASInt info;
995:   #if PetscDefined(USE_COMPLEX)
996:       PetscReal rwork_query;
997:       LAPACKsyev_("V", "U", &n_blas, &work_query, &n_blas, &rwork_query, &work_query, &lwork_query, &rwork_query, &info);
998:   #else
999:       LAPACKsyev_("V", "U", &n_blas, &work_query, &n_blas, &work_query, &work_query, &lwork_query, &info);
1000:   #endif
1001:       PetscCheck(info == 0, PETSC_COMM_SELF, PETSC_ERR_LIB, "LAPACK workspace query failed on SYEV %" PetscBLASInt_FMT, info);
1002:       eigen_work->lwork = (PetscBLASInt)PetscRealPart(work_query);

1004:       /* Allocate workspace */
1005:   #if PetscDefined(USE_COMPLEX)
1006:       PetscCall(PetscMalloc4(chunk_size * m * m, &eigen_work->all_v, chunk_size * m, &eigen_work->all_lambda, chunk_size * eigen_work->lwork, &eigen_work->all_work, chunk_size * (3 * m - 2), &eigen_work->all_rwork));
1007:   #else
1008:       PetscCall(PetscMalloc3(chunk_size * m * m, &eigen_work->all_v, chunk_size * m, &eigen_work->all_lambda, chunk_size * eigen_work->lwork, &eigen_work->all_work));
1009:   #endif
1010:     }
1011: #endif
1012:   }

1014:   /* Local aliases so KOKKOS_LAMBDAs capture views by value, not via eigen_work-> */
1015:   S_batch               = eigen_work->S_batch;
1016:   T_batch               = eigen_work->T_batch;
1017:   V_batch               = eigen_work->V_batch;
1018:   Lambda_batch          = eigen_work->Lambda_batch;
1019:   T_sqrt_batch          = eigen_work->T_sqrt_batch;
1020:   w_batch               = eigen_work->w_batch;
1021:   delta_batch           = eigen_work->delta_batch;
1022:   y_batch               = eigen_work->y_batch;
1023:   y_mean_batch          = eigen_work->y_mean_batch;
1024:   r_inv_sqrt_batch      = eigen_work->r_inv_sqrt_batch;
1025:   temp1_batch           = eigen_work->temp1_batch;
1026:   temp2_batch           = eigen_work->temp2_batch;
1027:   inv_sqrt_lambda_batch = eigen_work->inv_sqrt_lambda_batch;

1029:   /* Loop over chunks */
1030:   for (PetscInt chunk_start = 0; chunk_start < n_vertices; chunk_start += chunk_size) {
1031:     PetscInt chunk_end       = (chunk_start + chunk_size > n_vertices) ? n_vertices : chunk_start + chunk_size;
1032:     PetscInt n_batch_current = chunk_end - chunk_start;

1034:     /* No pre-zeroing of S_batch/delta_batch/y_*_batch/r_inv_sqrt_batch is required: the
1035:        fused extractor writes positions [0, ncols) on every iteration, and every downstream
1036:        consumer (ComputeAllTMatrices, ComputeWeightsAndInvSqrtLambda, the DEBUG NaN check)
1037:        reads through subviews bounded by the per-row ncols. Stale values in [ncols, max_nnz_per_row)
1038:        carried over from a previous chunk are unreachable. */

1040:     /* ===================================================================== */
1041:     /* Step 2.1.2: Fused observation extraction and S/Delta computation     */
1042:     /* ===================================================================== */
1043:     /* Extract local observations and immediately compute S and delta       */
1044:     /* This fusion eliminates one kernel launch and improves cache locality */
1045:     Kokkos::parallel_for(
1046:       "ExtractAndComputeSAndDelta", Kokkos::RangePolicy<exec_space>(0, n_batch_current), KOKKOS_LAMBDA(const int i_local) {
1047:         PetscInt i_global = chunk_start + i_local;
1048:         /* Get Q row for this grid point using CSR format */
1049:         PetscInt row_start = Q_i_view(i_global);
1050:         PetscInt row_end   = Q_i_view(i_global + 1);
1051:         PetscInt ncols     = row_end - row_start;

1053:         /* Extract observations and compute S/delta for this grid point */
1054:         for (PetscInt k = 0; k < ncols; k++) {
1055:           PetscInt    obs_idx = Q_j_view(row_start + k);
1056:           PetscScalar weight  = Q_a_view(row_start + k);

1058:           /* Extract observation vectors */
1059:           PetscScalar y_val      = y_global_view(obs_idx);
1060:           PetscScalar y_mean_val = y_mean_global_view(obs_idx);
1061:           PetscScalar r_inv_sqrt = r_inv_sqrt_global_view(obs_idx) * Kokkos::sqrt(PetscRealPart(weight));

1063:           /* Store for later use if needed */
1064:           y_batch(i_local, k)          = y_val;
1065:           y_mean_batch(i_local, k)     = y_mean_val;
1066:           r_inv_sqrt_batch(i_local, k) = r_inv_sqrt;

1068:           /* Compute delta immediately: delta = R^{-1/2}(y - y_mean) */
1069:           delta_batch(i_local, k) = (y_val - y_mean_val) * r_inv_sqrt;

1071:           /* Compute S row: S = R^{-1/2}(Z - y_mean * 1')/sqrt(m-1) */
1072:           PetscScalar scale_factor = scale * r_inv_sqrt;
1073:           for (int j = 0; j < m; j++) S_batch(i_local, k, j) = (Z_global_view(obs_idx, j) - y_mean_val) * scale_factor;
1074:         }
1075:       });
1076:     Kokkos::fence();

1078:     /* DEBUG: Check S for NaNs */
1079:     if (PetscDefined(USE_DEBUG)) {
1080:       PetscInt nan_count = 0;
1081:       Kokkos::parallel_reduce(
1082:         "CheckS", Kokkos::RangePolicy<exec_space>(0, n_batch_current),
1083:         KOKKOS_LAMBDA(const int i, PetscInt &l_count) {
1084:           PetscInt i_global = chunk_start + i;
1085:           PetscInt ncols    = Q_i_view(i_global + 1) - Q_i_view(i_global);
1086:           for (PetscInt j = 0; j < ncols; j++) {
1087:             for (int k = 0; k < m; k++) {
1088:               if (S_batch(i, j, k) != S_batch(i, j, k)) l_count++;
1089:             }
1090:           }
1091:         },
1092:         nan_count);
1093:       PetscCheck(nan_count == 0, PETSC_COMM_SELF, PETSC_ERR_FP, "Found %" PetscInt_FMT " NaNs in S_batch at chunk_start %" PetscInt_FMT, nan_count, chunk_start);
1094:     }

1096:     /* ===================================================================== */
1097:     /* Step 2.1.4: Optimized T matrix formation (T = (1/rho)I + S^T * S)    */
1098:     /* ===================================================================== */
1099:     /* Compute T_i = (1/rho)I + S_i^T * S_i for current chunk */
1100:     /* Exploit symmetry: only compute upper triangle, then copy to lower */
1101:     /* This reduces operations by ~50% */
1102:     Kokkos::parallel_for(
1103:       "ComputeAllTMatrices", Kokkos::RangePolicy<exec_space>(0, n_batch_current), KOKKOS_LAMBDA(const int i) {
1104:         PetscInt i_global = chunk_start + i;
1105:         PetscInt ncols    = Q_i_view(i_global + 1) - Q_i_view(i_global);

1107:         /* Compute upper triangle of T_i = (1/rho)I + S_i^T * S_i */
1108:         /* T_i(j,k) = (1/rho)*delta_jk + sum_p S_i(p,j) * S_i(p,k) for j <= k */
1109:         for (int j = 0; j < m; j++) {
1110:           for (int k = j; k < m; k++) {
1111:             PetscScalar sum = (j == k) ? inflation_inv : 0.0;
1112:             for (PetscInt p = 0; p < ncols; p++) sum += S_batch(i, p, j) * S_batch(i, p, k);
1113:             T_batch(i, j, k) = sum;
1114:           }
1115:         }

1117:         /* Copy upper triangle to lower triangle (T is symmetric) */
1118:         for (int j = 0; j < m; j++) {
1119:           for (int k = 0; k < j; k++) T_batch(i, j, k) = T_batch(i, k, j);
1120:         }
1121:       });
1122:     Kokkos::fence();

1124:     /* DEBUG: Check T for NaNs */
1125:     if (PetscDefined(USE_DEBUG)) {
1126:       PetscInt nan_count = 0;
1127:       Kokkos::parallel_reduce(
1128:         "CheckT", Kokkos::RangePolicy<exec_space>(0, n_batch_current),
1129:         KOKKOS_LAMBDA(const int i, PetscInt &l_count) {
1130:           for (int j = 0; j < m; j++) {
1131:             for (int k = 0; k < m; k++) {
1132:               if (T_batch(i, j, k) != T_batch(i, j, k)) l_count++;
1133:             }
1134:           }
1135:         },
1136:         nan_count);
1137:       PetscCheck(nan_count == 0, PETSC_COMM_SELF, PETSC_ERR_FP, "Found %" PetscInt_FMT " NaNs in T_batch at chunk_start %" PetscInt_FMT, nan_count, chunk_start);
1138:     }

1140:     /* ===================================================================== */
1141:     /* Step 3.1.1: Batched eigendecomposition for current chunk            */
1142:     /* ===================================================================== */
1143:     /* Compute T_i = V_i * Lambda_i * V_i^T for current chunk */
1144: #if defined(KOKKOS_ENABLE_CUDA) || defined(KOKKOS_ENABLE_HIP) || defined(KOKKOS_ENABLE_SYCL)
1145:     PetscCall(BatchedEigenSolve(T_batch, Lambda_batch, V_batch, n_batch_current, m, device_handle, eigen_work));
1146: #else
1147:     PetscCall(BatchedEigenSolve(T_batch, Lambda_batch, V_batch, n_batch_current, m, eigen_work));
1148: #endif

1150:     /* DEBUG: Check Lambda for NaNs or negative values */
1151:     if (PetscDefined(USE_DEBUG)) {
1152:       PetscInt bad_lambda = 0;
1153:       Kokkos::parallel_reduce(
1154:         "CheckLambda", Kokkos::RangePolicy<exec_space>(0, n_batch_current),
1155:         KOKKOS_LAMBDA(const int i, PetscInt &l_count) {
1156:           for (int k = 0; k < m; k++) {
1157:             if (Lambda_batch(i, k) != Lambda_batch(i, k) || PetscRealPart(Lambda_batch(i, k)) < -1e-8) l_count++;
1158:           }
1159:         },
1160:         bad_lambda);
1161:       PetscCheck(bad_lambda == 0, PETSC_COMM_SELF, PETSC_ERR_FP, "Found %" PetscInt_FMT " bad eigenvalues (NaN or negative) at chunk_start %" PetscInt_FMT, bad_lambda, chunk_start);
1162:     }

1164:     /* ===================================================================== */
1165:     /* Step 3.1.2: Precompute w and inv_sqrt_lambda for ensemble update    */
1166:     /* ===================================================================== */
1167:     /* Compute w_i = T_i^{-1} * (S_i^T * delta_i) using eigendecomposition */
1168:     /* Precompute 1/sqrt(Lambda) for use in ensemble update */
1169:     Kokkos::parallel_for(
1170:       "ComputeWeightsAndInvSqrtLambda", Kokkos::RangePolicy<exec_space>(0, n_batch_current), KOKKOS_LAMBDA(const int i) {
1171:         PetscInt i_global          = chunk_start + i;
1172:         PetscInt ncols             = Q_i_view(i_global + 1) - Q_i_view(i_global);
1173:         auto     S_i               = Kokkos::subview(S_batch, i, Kokkos::make_pair((PetscInt)0, ncols), Kokkos::ALL());
1174:         auto     V_i               = Kokkos::subview(V_batch, i, Kokkos::ALL(), Kokkos::ALL());
1175:         auto     Lambda_i          = Kokkos::subview(Lambda_batch, i, Kokkos::ALL());
1176:         auto     delta_i           = Kokkos::subview(delta_batch, i, Kokkos::make_pair((PetscInt)0, ncols));
1177:         auto     w_i               = Kokkos::subview(w_batch, i, Kokkos::ALL());
1178:         auto     inv_sqrt_lambda_i = Kokkos::subview(inv_sqrt_lambda_batch, i, Kokkos::ALL());
1179:         auto     temp1             = Kokkos::subview(temp1_batch, i, Kokkos::ALL());
1180:         auto     temp2             = Kokkos::subview(temp2_batch, i, Kokkos::ALL());

1182:         /* 1. Compute w_i = V * L^-1 * V^T * S^T * delta */
1183:         /* Step 1a: temp1 = S^T * delta using KokkosBlas::gemv for better vectorization */
1184:         KokkosBlas::SerialGemv<KokkosBlas::Trans::Transpose, KokkosBlas::Algo::Gemv::Unblocked>::invoke(1.0, S_i, delta_i, 0.0, temp1);

1186:         /* Step 1b: temp2 = V^T * temp1 using KokkosBlas::gemv for better vectorization */
1187:         KokkosBlas::SerialGemv<KokkosBlas::Trans::Transpose, KokkosBlas::Algo::Gemv::Unblocked>::invoke(1.0, V_i, temp1, 0.0, temp2);

1189:         /* Step 1c: temp2 = temp2 / Lambda; floor Lambda by LETKF_EIGEN_EPS (see header) */
1190:         for (int j = 0; j < m; j++) temp2(j) /= (Lambda_i(j) + LETKF_EIGEN_EPS);

1192:         /* Step 1d: w = V * temp2 using KokkosBlas::gemv for better vectorization */
1193:         KokkosBlas::SerialGemv<KokkosBlas::Trans::NoTranspose, KokkosBlas::Algo::Gemv::Unblocked>::invoke(1.0, V_i, temp2, 0.0, w_i);

1195:         /* 2. Precompute 1/sqrt(Lambda) for ensemble update; same LETKF_EIGEN_EPS floor as above */
1196:         for (int p = 0; p < m; p++) inv_sqrt_lambda_i(p) = 1.0 / Kokkos::sqrt(PetscRealPart(Lambda_i(p)) + LETKF_EIGEN_EPS);
1197:       });
1198:     Kokkos::fence();

1200:     /* ===================================================================== */
1201:     /* Step 3.1.3: Fused G computation and ensemble update                  */
1202:     /* ===================================================================== */
1203:     /* Compute E[i,:] = mean[i] + X[i,:] * G_i on-the-fly */
1204:     /* G_i is computed column-by-column and immediately applied */
1205:     /* This eliminates the need to store G_batch, saving m*m*n_batch memory */
1206:     Kokkos::parallel_for(
1207:       "FusedGComputeAndEnsembleUpdate", Kokkos::RangePolicy<exec_space>(0, n_batch_current), KOKKOS_LAMBDA(const int i_local) {
1208:         PetscInt i_global = chunk_start + i_local;

1210:         auto X_i    = Kokkos::subview(X_view, Kokkos::make_pair(i_global * ndof, (i_global + 1) * ndof), Kokkos::ALL());
1211:         auto E_i    = Kokkos::subview(E_view, Kokkos::make_pair(i_global * ndof, (i_global + 1) * ndof), Kokkos::ALL());
1212:         auto mean_i = Kokkos::subview(mean_view, Kokkos::make_pair(i_global * ndof, (i_global + 1) * ndof));

1214:         auto V_i               = Kokkos::subview(V_batch, i_local, Kokkos::ALL(), Kokkos::ALL());
1215:         auto w_i               = Kokkos::subview(w_batch, i_local, Kokkos::ALL());
1216:         auto inv_sqrt_lambda_i = Kokkos::subview(inv_sqrt_lambda_batch, i_local, Kokkos::ALL());
1217:         auto T_sqrt_i          = Kokkos::subview(T_sqrt_batch, i_local, Kokkos::ALL(), Kokkos::ALL());

1219:         /* Initialize E_i with mean */
1220:         for (int row = 0; row < ndof; row++) {
1221:           PetscScalar m_val = mean_i(row);
1222:           for (int col = 0; col < m; col++) E_i(row, col) = m_val;
1223:         }

1225:         /* Compute T_sqrt = V * diag(1/sqrt(Lambda)) * V^T */
1226:         /* Optimized: Exploit symmetry - only compute upper triangle, then copy to lower */
1227:         /* T_sqrt(j,k) = sum_p V(j,p) * V(k,p) / sqrt(Lambda(p)) for j <= k */
1228:         for (int j = 0; j < m; j++) {
1229:           for (int k = j; k < m; k++) {
1230:             PetscScalar sum = 0.0;
1231:             for (int p = 0; p < m; p++) sum += V_i(j, p) * V_i(k, p) * inv_sqrt_lambda_i(p);
1232:             T_sqrt_i(j, k) = sum;
1233:           }
1234:         }
1235:         /* Copy upper triangle to lower triangle (T_sqrt is symmetric) */
1236:         for (int j = 0; j < m; j++) {
1237:           for (int k = 0; k < j; k++) T_sqrt_i(j, k) = T_sqrt_i(k, j);
1238:         }

1240:         /* Compute E_i += X_i * G_i column-by-column */
1241:         /* G_i(:,k) = w_i + sqrt(m-1) * T_sqrt_i(:,k) */
1242:         for (int k = 0; k < m; k++) {
1243:           /* Compute column k of G on-the-fly */
1244:           for (int row = 0; row < ndof; row++) {
1245:             PetscScalar sum = 0.0;
1246:             for (int j = 0; j < m; j++) {
1247:               /* G_i(j,k) = w_i(j) + sqrt(m-1) * T_sqrt_i(j,k) */
1248:               PetscScalar G_jk = w_i(j) + sqrt_m_minus_1 * T_sqrt_i(j, k);
1249:               sum += X_i(row, j) * G_jk;
1250:             }
1251:             E_i(row, k) += sum;
1252:           }
1253:         }
1254:       });
1255:     Kokkos::fence();
1256:   }

1258:   /* Cleanup workspace */
1259:   /* NOTE: Workspace is now persistent in impl->eigen_work and impl->solver_handle */
1260:   /* It will be destroyed in PetscDALETKFDestroyLocalization_Kokkos */

1262:   /* Copy back updated ensemble if needed */
1263:   if (e_is_copy) {
1264:     Kokkos::View<PetscScalar **, Kokkos::LayoutLeft, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>> dst(e_array, lda_e, m);
1265:     PetscCallCXX(Kokkos::deep_copy(dst, e_managed));
1266:   }

1268:   /* Restore arrays */
1269:   PetscCall(MatDenseRestoreArrayWriteAndMemType(en->ensemble, &e_array));
1270:   PetscCall(VecRestoreArrayReadAndMemType(impl->mean, &mean_array));
1271:   PetscCall(MatDenseRestoreArrayReadAndMemType(X, &x_array));

1273:   /* Restore global observation arrays */
1274:   PetscCall(VecRestoreArrayReadAndMemType(r_inv_sqrt_global, &r_inv_sqrt_global_array));
1275:   PetscCall(VecRestoreArrayReadAndMemType(y_mean_global, &y_mean_global_array));
1276:   PetscCall(VecRestoreArrayReadAndMemType(observation, &y_global_array));
1277:   PetscCall(MatDenseRestoreArrayReadAndMemType(Z_global, &z_global_array));

1279:   /* impl->n_nnz_local was populated by PetscDALETKFInstallQ() and is valid for the lifetime of Q
1280:      (which only changes via InstallQ). Reading from impl avoids a redundant MatGetInfo here (and
1281:      the AIJKOKKOS device->host sync it would trigger). */
1282:   n_obs_total = (PetscReal)impl->n_nnz_local;
1283:   flops       = 0.0;

1285:   /* Step 2.1.2: Fused observation extraction and S/Delta computation */
1286:   flops += n_obs_total * (2.0 + 2.0 * m);

1288:   /* Step 2.1.4: Optimized T matrix formation */
1289:   flops += n_obs_total * m * (m + 1);

1291:   /* Step 3.1.2: Precompute w and inv_sqrt_lambda */
1292:   flops += n_obs_total * 2.0 * m + (PetscReal)n_vertices * (4.0 * m * m + 3.0 * m);

1294:   /* Step 3.1.3: Fused G computation and ensemble update */
1295:   /* T_sqrt: 1.5*m^3 + 1.5*m^2 */
1296:   flops += (PetscReal)n_vertices * (1.5 * m * m * m + 1.5 * m * m);
1297:   /* E update: ndof * m * (4*m + 1) */
1298:   /* Note: G_jk computation (2 flops) is inside the inner loop, so it's 2*m*ndof*m */
1299:   /* Matrix product X*G (2 flops) is also 2*m*ndof*m */
1300:   flops += (PetscReal)n_vertices * ndof * m * (4.0 * m + 1.0);

1302:   PetscCall(PetscLogGpuFlops(flops));
1303:   PetscFunctionReturn(PETSC_SUCCESS);
1304: }

1306: /* ========================================================================== */
1307: /*                LETKF Global Analysis (LOC_NONE, Kokkos path)               */
1308: /* ========================================================================== */

1310: /*
1311:   PetscDALETKFGlobalAnalysis_Kokkos - LOC_NONE LETKF analysis using device gemm/gemv.

1313:   Mirrors the CPU LOC_NONE block in PetscDAEnsembleAnalysis_LETKF: device-side gemm for
1314:   S^T*S, gemv for S^T*delta, and gemm for X*G; the m x m factor (T = (1/rho)I + S^T*S)
1315:   is eigendecomposed on the host on every rank since m is small (ensemble size).
1316: */
1317: PETSC_INTERN PetscErrorCode PetscDALETKFGlobalAnalysis_Kokkos(PetscDA da, PetscDA_LETKF *impl, PetscInt m, Mat X, Vec observation)
1318: {
1319:   using exec_space       = Kokkos::DefaultExecutionSpace;
1320:   using view_2d          = Kokkos::View<PetscScalar **, Kokkos::LayoutLeft, exec_space>;
1321:   using view_1d          = Kokkos::View<PetscScalar *, Kokkos::LayoutLeft, exec_space>;
1322:   using view_2d_const_um = Kokkos::View<const PetscScalar **, Kokkos::LayoutLeft, exec_space, Kokkos::MemoryTraits<Kokkos::Unmanaged>>;
1323:   using view_1d_const_um = Kokkos::View<const PetscScalar *, Kokkos::LayoutLeft, exec_space, Kokkos::MemoryTraits<Kokkos::Unmanaged>>;
1324:   using view_2d_um       = Kokkos::View<PetscScalar **, Kokkos::LayoutLeft, exec_space, Kokkos::MemoryTraits<Kokkos::Unmanaged>>;
1325:   using h_2d_const_um    = Kokkos::View<const PetscScalar **, Kokkos::LayoutLeft, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>>;
1326:   using h_1d_const_um    = Kokkos::View<const PetscScalar *, Kokkos::LayoutLeft, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>>;
1327:   using h_2d_um          = Kokkos::View<PetscScalar **, Kokkos::LayoutLeft, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>>;
1328:   using h_1d_um          = Kokkos::View<PetscScalar *, Kokkos::LayoutLeft, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>>;
1329:   MPI_Comm           comm;
1330:   PetscReal          scale, sqrt_m_minus_1;
1331:   PetscInt           n_obs_local, n_local_ens, s_lda, x_lda, e_lda, g_lda;
1332:   PetscMemType       s_mt, d_mt, mean_mt, x_mt, e_mt;
1333:   const PetscScalar *s_arr, *d_arr, *mean_arr, *x_arr, *g_host;
1334:   const PetscScalar *s_dev, *d_dev, *mean_dev, *x_dev;
1335:   PetscScalar       *e_arr, *e_dev, *gram_host, *sd_buf;
1336:   PetscMPIInt        mMPI, mmMPI;
1337:   PetscBool          e_is_copy = PETSC_FALSE;
1338:   view_2d            S_managed, X_managed, E_managed, gram_dev, G_dev, XG_dev;
1339:   view_1d            d_managed, mean_managed, Sd_dev;

1341:   PetscFunctionBegin;
1342:   PetscCall(PetscObjectGetComm((PetscObject)da, &comm));
1343:   PetscCall(PetscKokkosInitializeCheck());

1345:   scale          = 1.0 / PetscSqrtReal((PetscReal)(m - 1));
1346:   sqrt_m_minus_1 = PetscSqrtReal((PetscReal)(m - 1));

1348:   PetscCall(PetscDALETKFEnsureGlobalScratch(impl, m));

1350:   /* S = R^{-1/2} * (Z - y_mean*1') / sqrt(m-1); delta_scaled = R^{-1/2} * (y - y_mean) */
1351:   PetscCall(PetscDAEnsembleComputeNormalizedInnovationMatrix(impl->Z, impl->y_mean, impl->r_inv_sqrt, m, scale, impl->S));
1352:   PetscCall(VecWAXPY(impl->delta_scaled, -1.0, impl->y_mean, observation));
1353:   PetscCall(VecPointwiseMult(impl->delta_scaled, impl->delta_scaled, impl->r_inv_sqrt));

1355:   PetscCall(MatGetLocalSize(impl->S, &n_obs_local, NULL));
1356:   PetscCall(MatGetLocalSize(impl->en.ensemble, &n_local_ens, NULL));

1358:   PetscCall(MatDenseGetArrayReadAndMemType(impl->S, &s_arr, &s_mt));
1359:   PetscCall(MatDenseGetLDA(impl->S, &s_lda));
1360:   PetscCall(VecGetArrayReadAndMemType(impl->delta_scaled, &d_arr, &d_mt));
1361:   PetscCall(VecGetArrayReadAndMemType(impl->mean, &mean_arr, &mean_mt));
1362:   PetscCall(MatDenseGetArrayReadAndMemType(X, &x_arr, &x_mt));
1363:   PetscCall(MatDenseGetLDA(X, &x_lda));
1364:   PetscCall(MatDenseGetArrayWriteAndMemType(impl->en.ensemble, &e_arr, &e_mt));
1365:   PetscCall(MatDenseGetLDA(impl->en.ensemble, &e_lda));

1367:   /* Mirror host arrays to device when needed. */
1368:   s_dev    = s_arr;
1369:   d_dev    = d_arr;
1370:   mean_dev = mean_arr;
1371:   x_dev    = x_arr;
1372:   e_dev    = e_arr;

1374:   if (s_mt == PETSC_MEMTYPE_HOST && n_obs_local > 0) {
1375:     PetscCallCXX(S_managed = view_2d("S_managed", s_lda, m));
1376:     PetscCallCXX(Kokkos::deep_copy(S_managed, h_2d_const_um(s_arr, s_lda, m)));
1377:     s_dev = S_managed.data();
1378:   }
1379:   if (d_mt == PETSC_MEMTYPE_HOST && n_obs_local > 0) {
1380:     PetscCallCXX(d_managed = view_1d("d_managed", n_obs_local));
1381:     PetscCallCXX(Kokkos::deep_copy(d_managed, h_1d_const_um(d_arr, n_obs_local)));
1382:     d_dev = d_managed.data();
1383:   }
1384:   if (mean_mt == PETSC_MEMTYPE_HOST && n_local_ens > 0) {
1385:     PetscCallCXX(mean_managed = view_1d("mean_managed", n_local_ens));
1386:     PetscCallCXX(Kokkos::deep_copy(mean_managed, h_1d_const_um(mean_arr, n_local_ens)));
1387:     mean_dev = mean_managed.data();
1388:   }
1389:   /* X_managed / E_managed are only consumed inside the n_local_ens > 0 block below; allocating
1390:      the mirrors on a rank with no ensemble rows would size to a non-positive (x_lda, e_lda) and
1391:      waste a Kokkos View on data we never read or write. */
1392:   if (x_mt == PETSC_MEMTYPE_HOST && n_local_ens > 0) {
1393:     PetscCallCXX(X_managed = view_2d("X_managed", x_lda, m));
1394:     PetscCallCXX(Kokkos::deep_copy(X_managed, h_2d_const_um(x_arr, x_lda, m)));
1395:     x_dev = X_managed.data();
1396:   }
1397:   if (e_mt == PETSC_MEMTYPE_HOST && n_local_ens > 0) {
1398:     PetscCallCXX(E_managed = view_2d("E_managed", e_lda, m));
1399:     e_dev     = E_managed.data();
1400:     e_is_copy = PETSC_TRUE;
1401:   }

1403:   /* Device gemm: gram = S^T * S over the active local rows [0, n_obs_local). */
1404:   PetscCallCXX(gram_dev = view_2d("gram_dev", m, m));
1405:   if (n_obs_local > 0) {
1406:     view_2d_const_um S_full(s_dev, s_lda, m);
1407:     auto             S_active = Kokkos::subview(S_full, Kokkos::make_pair((PetscInt)0, n_obs_local), Kokkos::ALL());
1408:     KokkosBlas::gemm("T", "N", (PetscScalar)1.0, S_active, S_active, (PetscScalar)0.0, gram_dev);
1409:   }
1410:   Kokkos::fence();

1412:   /* Mirror gram to host, allreduce, and feed the shared SELF-gram factorizer. PetscCalloc1
1413:      so an n_obs_local == 0 rank (where the device gemm above is skipped, leaving gram_dev
1414:      at Kokkos's default-zero state but a future refactor might also skip the deep_copy)
1415:      contributes zeros to the allreduce instead of uninitialized bytes. */
1416:   PetscCall(PetscCalloc1((size_t)m * m, &gram_host));
1417:   PetscCallCXX(Kokkos::deep_copy(h_2d_um(gram_host, m, m), gram_dev));
1418:   PetscCall(PetscMPIIntCast((PetscInt64)m * m, &mmMPI));
1419:   PetscCallMPI(MPIU_Allreduce(MPI_IN_PLACE, gram_host, mmMPI, MPIU_SCALAR, MPIU_SUM, comm));
1420:   PetscCall(PetscDAEnsembleTFactorFromGram(da, m, gram_host));
1421:   PetscCall(PetscFree(gram_host));

1423:   /* Device gemv: Sd = S^T * delta_scaled, then mirror + Allreduce. */
1424:   PetscCallCXX(Sd_dev = view_1d("Sd_dev", m));
1425:   if (n_obs_local > 0) {
1426:     view_2d_const_um S_full(s_dev, s_lda, m);
1427:     view_1d_const_um d_full(d_dev, n_obs_local);
1428:     auto             S_active = Kokkos::subview(S_full, Kokkos::make_pair((PetscInt)0, n_obs_local), Kokkos::ALL());
1429:     KokkosBlas::gemv("T", (PetscScalar)1.0, S_active, d_full, (PetscScalar)0.0, Sd_dev);
1430:   }
1431:   Kokkos::fence();

1433:   /* Stage the device-side Sd into the persistent impl->s_transpose_delta scratch (matches the
1434:      CPU path) and allreduce in place. */
1435:   PetscCall(VecGetArray(impl->s_transpose_delta, &sd_buf));
1436:   PetscCallCXX(Kokkos::deep_copy(h_1d_um(sd_buf, m), Sd_dev));
1437:   PetscCall(PetscMPIIntCast(m, &mMPI));
1438:   PetscCallMPI(MPIU_Allreduce(MPI_IN_PLACE, sd_buf, mMPI, MPIU_SCALAR, MPIU_SUM, comm));
1439:   PetscCall(VecRestoreArray(impl->s_transpose_delta, &sd_buf));

1441:   /* Restore S/delta read arrays before host-side T-inverse. */
1442:   PetscCall(VecRestoreArrayReadAndMemType(impl->delta_scaled, &d_arr));
1443:   PetscCall(MatDenseRestoreArrayReadAndMemType(impl->S, &s_arr));

1445:   /* w = T^{-1} * (S^T * delta), all on PETSC_COMM_SELF. */
1446:   PetscCall(PetscDAEnsembleApplyTInverse(da, impl->s_transpose_delta, impl->w));

1448:   /* T_sqrt = T^{-1/2} on PETSC_COMM_SELF. */
1449:   PetscCall(PetscDAEnsembleApplySqrtTInverse(da, NULL, impl->T_sqrt));

1451:   /* G = w*1' + sqrt(m-1) * T_sqrt, m x m on PETSC_COMM_SELF in impl->w_ones. */
1452:   PetscCall(PetscDALETKFReplicateWeightVector(impl->w, m, impl->w_ones));
1453:   PetscCall(MatAXPY(impl->w_ones, sqrt_m_minus_1, impl->T_sqrt, SAME_NONZERO_PATTERN));

1455:   /* Push G to device for the X*G gemm. impl->w_ones is a SELF SeqDense; LDA == m. */
1456:   PetscCallCXX(G_dev = view_2d("G_dev", m, m));
1457:   PetscCall(MatDenseGetArrayRead(impl->w_ones, &g_host));
1458:   PetscCall(MatDenseGetLDA(impl->w_ones, &g_lda));
1459:   PetscCheck(g_lda == m, PETSC_COMM_SELF, PETSC_ERR_PLIB, "Unexpected LDA %" PetscInt_FMT " for SELF SeqDense w_ones (m=%" PetscInt_FMT ")", g_lda, m);
1460:   PetscCallCXX(Kokkos::deep_copy(G_dev, h_2d_const_um(g_host, m, m)));
1461:   PetscCall(MatDenseRestoreArrayRead(impl->w_ones, &g_host));

1463:   /* Device gemm: XG = X_local * G, then E = mean*1' + XG. Allocate XG_dev only when
1464:      this rank actually owns ensemble columns; consumers below are gated identically. */
1465:   if (n_local_ens > 0) {
1466:     view_2d_const_um X_full(x_dev, x_lda, m);
1467:     auto             X_active = Kokkos::subview(X_full, Kokkos::make_pair((PetscInt)0, n_local_ens), Kokkos::ALL());
1468:     view_2d_um       E_full(e_dev, e_lda, m);
1469:     view_1d_const_um mean_full(mean_dev, n_local_ens);
1470:     PetscInt         m_local = m;

1472:     PetscCallCXX(XG_dev = view_2d("XG_dev", n_local_ens, m));
1473:     KokkosBlas::gemm("N", "N", (PetscScalar)1.0, X_active, G_dev, (PetscScalar)0.0, XG_dev);
1474:     Kokkos::parallel_for(
1475:       "EnsembleUpdate_LOC_NONE", Kokkos::RangePolicy<exec_space>(0, n_local_ens), KOKKOS_LAMBDA(const int i) {
1476:         PetscScalar mi = mean_full(i);
1477:         for (PetscInt j = 0; j < m_local; j++) E_full(i, j) = mi + XG_dev(i, j);
1478:       });
1479:     Kokkos::fence();
1480:   }

1482:   PetscCall(VecRestoreArrayReadAndMemType(impl->mean, &mean_arr));
1483:   PetscCall(MatDenseRestoreArrayReadAndMemType(X, &x_arr));

1485:   if (e_is_copy && n_local_ens > 0) {
1486:     /* Copy only the active rows [0, n_local_ens). E_managed is allocated to (e_lda, m) so the
1487:        device-side strides line up with the host MatDense LDA, but the LDA-padding rows
1488:        [n_local_ens, e_lda) are not touched by the analysis and must not be written back into
1489:        the host buffer's opaque padding bytes. */
1490:     Kokkos::View<PetscScalar **, Kokkos::LayoutLeft, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>> dst(e_arr, e_lda, m);
1491:     PetscCallCXX(Kokkos::deep_copy(Kokkos::subview(dst, Kokkos::make_pair((PetscInt)0, n_local_ens), Kokkos::ALL()), Kokkos::subview(E_managed, Kokkos::make_pair((PetscInt)0, n_local_ens), Kokkos::ALL())));
1492:   }
1493:   PetscCall(MatDenseRestoreArrayWriteAndMemType(impl->en.ensemble, &e_arr));
1494:   PetscFunctionReturn(PETSC_SUCCESS);
1495: }