Actual source code: letkf_kernels.h

  1: #pragma once

  3: /*
  4:   Distance-based localization kernels shared by the CPU and Kokkos backends of LETKF.
  5:   All functions are pure (no PETSc state), so the same definitions are used host-only
  6:   from `dalocalizationletkf.c` and host+device from `dalocalizationletkf.kokkos.cxx`.

  8:   When compiled in a TU that has already included Kokkos, `KOKKOS_INLINE_FUNCTION` is
  9:   defined and the functions are annotated for both host and device. Otherwise they fall
 10:   back to plain `static inline`.

 12:   Include `<Kokkos_Core.hpp>` (or any header that pulls it in) BEFORE this header to get the
 13:   device-callable variants; otherwise the host-only fallbacks are silently selected.

 15:   Radius semantics differ by kernel; the user-supplied `radius` always controls the
 16:   effective support but the corresponding cutoff distance varies:
 17:     - Gaspari-Cohn: compactly supported with cutoff at distance = 2*radius
 18:     - Gaussian:     truncated at distance = 2*radius (value ~ exp(-2) ~ 0.135)
 19:     - Boxcar:       cutoff at distance = radius (1 inside, 0 outside)
 20:   Set `radius` per kernel rather than expecting identical falloff across choices.
 21: */

 23: #include <petsc.h>
 24: #include <petscda.h>

 26: #if defined(KOKKOS_INLINE_FUNCTION)
 27:   #define LETKF_KERNEL_FN               KOKKOS_INLINE_FUNCTION
 28:   #define LETKF_KERNEL_UNREACHABLE(...) Kokkos::abort(__VA_ARGS__)
 29:   #define LETKF_KERNEL_EXP(x)           Kokkos::exp(x)  /* device-callable; PetscExpReal would resolve to host libm */
 30:   #define LETKF_KERNEL_SQRT(x)          Kokkos::sqrt(x) /* device-callable; PetscSqrtReal would resolve to host libm */
 31: #else
 32:   #define LETKF_KERNEL_FN               static inline
 33:   #define LETKF_KERNEL_UNREACHABLE(...) SETERRABORT(PETSC_COMM_SELF, PETSC_ERR_PLIB, __VA_ARGS__)
 34:   #define LETKF_KERNEL_EXP(x)           PetscExpReal(x)
 35:   #define LETKF_KERNEL_SQRT(x)          PetscSqrtReal(x)
 36: #endif

 38: /* Gaspari-Cohn 5th-order piecewise polynomial. */
 39: LETKF_KERNEL_FN PetscReal LETKFGaspariCohn(PetscReal distance, PetscReal radius)
 40: {
 41:   PetscReal r, r2, r3, r4, r5, val;

 43:   if (radius <= 0.0) return 0.0;
 44:   r = distance / radius;
 45:   if (r >= 2.0) return 0.0;

 47:   r2 = r * r;
 48:   r3 = r2 * r;
 49:   r4 = r3 * r;
 50:   r5 = r4 * r;

 52:   if (r <= 1.0) val = -0.25 * r5 + 0.5 * r4 + 0.625 * r3 - (5.0 / 3.0) * r2 + 1.0;
 53:   else val = (1.0 / 12.0) * r5 - 0.5 * r4 + 0.625 * r3 + (5.0 / 3.0) * r2 - 5.0 * r + 4.0 - (2.0 / 3.0) / r;
 54:   return val;
 55: }

 57: /* Gaussian kernel exp(-d^2 / (2 r^2)), truncated at d = 2 r (value ~ exp(-2) ~ 0.135). */
 58: LETKF_KERNEL_FN PetscReal LETKFGaussian(PetscReal distance, PetscReal radius)
 59: {
 60:   PetscReal r;

 62:   if (radius <= 0.0) return 0.0;
 63:   r = distance / radius;
 64:   if (r >= 2.0) return 0.0;
 65:   return LETKF_KERNEL_EXP(-0.5 * r * r);
 66: }

 68: /* Boxcar kernel: 1 inside the radius, 0 outside. */
 69: LETKF_KERNEL_FN PetscReal LETKFBoxcar(PetscReal distance, PetscReal radius)
 70: {
 71:   if (radius <= 0.0) return 0.0;
 72:   return distance < radius ? 1.0 : 0.0;
 73: }

 75: /* Cutoff distance beyond which the kernel is guaranteed to return zero. Callers square
 76:    the result inline rather than going through a sqrt(cutoff^2) round-trip so the bbox
 77:    prune in PetscDALETKFGatherObsBbox() does not pick up a 1-ulp shrink that could drop
 78:    a boundary observation under the BOXCAR kernel's strict (distance < radius) test. */
 79: LETKF_KERNEL_FN PetscReal LETKFCutoff(PetscDALETKFLocalizationType type, PetscReal radius)
 80: {
 81:   switch (type) {
 82:   case PETSCDA_LETKF_LOC_GASPARI_COHN:
 83:   case PETSCDA_LETKF_LOC_GAUSSIAN:
 84:     return 2.0 * radius;
 85:   case PETSCDA_LETKF_LOC_BOXCAR:
 86:     return radius;
 87:   case PETSCDA_LETKF_LOC_NONE:
 88:   case PETSCDA_LETKF_LOC_NUM_TYPES:
 89:     break;
 90:   }
 91:   /* Unreachable: callers PetscCheck() the type before reaching here, and LOC_NONE skips Q construction entirely. */
 92:   LETKF_KERNEL_UNREACHABLE("LETKFCutoff: invalid localization type");
 93:   return 0.0;
 94: }

 96: LETKF_KERNEL_FN PetscReal LETKFKernelEval(PetscDALETKFLocalizationType type, PetscReal distance, PetscReal radius)
 97: {
 98:   switch (type) {
 99:   case PETSCDA_LETKF_LOC_GASPARI_COHN:
100:     return LETKFGaspariCohn(distance, radius);
101:   case PETSCDA_LETKF_LOC_GAUSSIAN:
102:     return LETKFGaussian(distance, radius);
103:   case PETSCDA_LETKF_LOC_BOXCAR:
104:     return LETKFBoxcar(distance, radius);
105:   case PETSCDA_LETKF_LOC_NONE:
106:   case PETSCDA_LETKF_LOC_NUM_TYPES:
107:     break;
108:   }
109:   /* Unreachable: see LETKFCutoff(). */
110:   LETKF_KERNEL_UNREACHABLE("LETKFKernelEval: invalid localization type");
111:   return 0.0;
112: }

114: /* Squared distance between two coordinate tuples with per-dimension minimum-image periodicity.
115:    For each d with bd[d] > 0 the difference is folded into [-bd[d]/2, bd[d]/2]; non-periodic dims
116:    pass through. The shape is identical across both passes of both backends, so factor it out. */
117: LETKF_KERNEL_FN PetscReal LETKFPeriodicDist2(PetscInt dim, const PetscReal *v, const PetscReal *o, const PetscReal *bd)
118: {
119:   PetscReal dist2 = 0.0;
120:   for (PetscInt d = 0; d < dim; ++d) {
121:     PetscReal diff = v[d] - o[d];
122:     if (bd[d] > 0.0) {
123:       PetscReal L = bd[d];
124:       if (diff > 0.5 * L) diff -= L;
125:       else if (diff < -0.5 * L) diff += L;
126:     }
127:     dist2 += diff * diff;
128:   }
129:   return dist2;
130: }

132: /* Localization weight for a single (vertex, observation) pair: zero outside the cutoff or when the
133:    kernel itself returns zero, else the kernel value. Encapsulates the dist2 + cutoff2 + KernelEval
134:    triple shared by both passes of the AIJ and Kokkos Q-construction loops. */
135: LETKF_KERNEL_FN PetscReal LETKFRowWeight(PetscDALETKFLocalizationType type, PetscReal radius, PetscReal cutoff2, PetscInt dim, const PetscReal *v, const PetscReal *o, const PetscReal *bd)
136: {
137:   PetscReal dist2 = LETKFPeriodicDist2(dim, v, o, bd);
138:   if (dist2 >= cutoff2) return 0.0;
139:   return LETKFKernelEval(type, LETKF_KERNEL_SQRT(dist2), radius);
140: }