@njit(parallel=True)
def _lw_transport_kernel(
tau, planck_source, surface_source, emissivity, weights,
up_band, down_band, up_broad, down_broad,
diag_trans, diag_up_gpt, diag_dn_gpt, want_diag, diffusivity_factor,
):
"""Consolidated multi-band, multi-g-point LW transport.
Loops over columns in parallel; for each (band, g-point) runs the up/down
diffusivity sweeps and accumulates weighted fluxes into up_band/down_band
inside the compiled kernel. Accumulation order (g ascending, then b
ascending for broadband) matches the original python loops bit-for-bit.
"""
nband, ngpt, nlev, ncol = tau.shape
for i in prange(ncol):
for k in range(nlev + 1):
up_broad[k, i] = 0.0
down_broad[k, i] = 0.0
for b in range(nband):
for k in range(nlev + 1):
up_band[b, k, i] = 0.0
down_band[b, k, i] = 0.0
for g in range(ngpt):
w = weights[b, g]
# Upward sweep: surface -> TOA
up_prev = emissivity[b, i] * surface_source[b, g, i]
up_band[b, 0, i] += w * up_prev
if want_diag != 0:
diag_up_gpt[b, g, 0, i] = w * up_prev
for k in range(nlev):
trans = np.exp(-diffusivity_factor * tau[b, g, k, i])
up_cur = up_prev * trans + planck_source[b, g, k, i] * (1.0 - trans)
up_band[b, k + 1, i] += w * up_cur
if want_diag != 0:
diag_trans[b, g, k, i] = trans
diag_up_gpt[b, g, k + 1, i] = w * up_cur
up_prev = up_cur
# Downward sweep: TOA -> surface (dn_prev starts at 0 = TOA BC)
dn_prev = 0.0
if want_diag != 0:
diag_dn_gpt[b, g, nlev, i] = 0.0
for k in range(nlev - 1, -1, -1):
trans = np.exp(-diffusivity_factor * tau[b, g, k, i])
dn_cur = dn_prev * trans + planck_source[b, g, k, i] * (1.0 - trans)
down_band[b, k, i] += w * dn_cur
if want_diag != 0:
diag_dn_gpt[b, g, k, i] = w * dn_cur
dn_prev = dn_cur
for k in range(nlev + 1):
up_broad[k, i] += up_band[b, k, i]
down_broad[k, i] += down_band[b, k, i]