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