GPU: fuse the whole RK4 discharge step into one cupy kernel
The naive elementwise port launched ~60 tiny kernels per step, so the GTX 1070 was only ~1.3x over CPU (launch-bound). Fold the entire RK4 step (4 derivative evals + combine + overlap) into a single cupy.fuse kernel; numpy path unchanged and still validates vs scipy. GPU correctness re-checked against CPU on deploy. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
@@ -109,6 +109,58 @@ def _derivatives(xp, q, i, x, v, p: BatchDischargeParams):
|
|||||||
return d_q, d_i, d_x, d_v
|
return d_q, d_i, d_x, d_v
|
||||||
|
|
||||||
|
|
||||||
|
_FUSED_STEP_CACHE = {}
|
||||||
|
|
||||||
|
|
||||||
|
def _get_fused_step(cp):
|
||||||
|
"""Собирает (и кэширует) cupy.fuse-ядро полного RK4-шага разряда.
|
||||||
|
|
||||||
|
Весь шаг (4 вычисления производных + сборка + overlap) сливается в ОДНО
|
||||||
|
GPU-ядро вместо ~60 мелких — убирает накладные на запуск ядер, из-за
|
||||||
|
которых наивный порт был лишь ~1.3x к CPU.
|
||||||
|
"""
|
||||||
|
if "step" in _FUSED_STEP_CACHE:
|
||||||
|
return _FUSED_STEP_CACHE["step"]
|
||||||
|
LN2 = 0.6931471805599453
|
||||||
|
|
||||||
|
def _ovl(x, hs, w):
|
||||||
|
s1 = 0.5 * (1.0 + cp.tanh((x + hs) / w * 0.5))
|
||||||
|
s2 = 0.5 * (1.0 + cp.tanh((hs - x) / w * 0.5))
|
||||||
|
return s1 * s2, (s1 * s2 / w) * (s2 - s1)
|
||||||
|
|
||||||
|
def _der(q, i, x, v, hs, w, isat, lair, liron, C, rt, red, m):
|
||||||
|
ov, dov = _ovl(x, hs, w)
|
||||||
|
z = i / isat
|
||||||
|
tz = cp.tanh(z)
|
||||||
|
g = isat * tz
|
||||||
|
gp = 1.0 - tz * tz
|
||||||
|
az = cp.abs(z)
|
||||||
|
G = isat * isat * (az + cp.log1p(cp.exp(-2.0 * az)) - LN2)
|
||||||
|
dl_di = lair + liron * ov * gp
|
||||||
|
reff = rt + red * ov
|
||||||
|
dq = -i
|
||||||
|
di = (q / C - i * reff - liron * dov * g * v) / dl_di
|
||||||
|
dv = (liron * dov * G) / m
|
||||||
|
return dq, di, v, dv
|
||||||
|
|
||||||
|
@cp.fuse()
|
||||||
|
def step(q, i, x, v, dt, hs, w, isat, lair, liron, C, rt, red, m):
|
||||||
|
a = _der(q, i, x, v, hs, w, isat, lair, liron, C, rt, red, m)
|
||||||
|
b = _der(q + dt * 0.5 * a[0], i + dt * 0.5 * a[1], x + dt * 0.5 * a[2], v + dt * 0.5 * a[3], hs, w, isat, lair, liron, C, rt, red, m)
|
||||||
|
c = _der(q + dt * 0.5 * b[0], i + dt * 0.5 * b[1], x + dt * 0.5 * b[2], v + dt * 0.5 * b[3], hs, w, isat, lair, liron, C, rt, red, m)
|
||||||
|
d = _der(q + dt * c[0], i + dt * c[1], x + dt * c[2], v + dt * c[3], hs, w, isat, lair, liron, C, rt, red, m)
|
||||||
|
qn = q + dt / 6.0 * (a[0] + 2 * b[0] + 2 * c[0] + d[0])
|
||||||
|
in_ = i + dt / 6.0 * (a[1] + 2 * b[1] + 2 * c[1] + d[1])
|
||||||
|
xn = x + dt / 6.0 * (a[2] + 2 * b[2] + 2 * c[2] + d[2])
|
||||||
|
vn = v + dt / 6.0 * (a[3] + 2 * b[3] + 2 * c[3] + d[3])
|
||||||
|
ov_old, _ = _ovl(x, hs, w)
|
||||||
|
ov_new, _ = _ovl(xn, hs, w)
|
||||||
|
return qn, in_, xn, vn, ov_old, ov_new
|
||||||
|
|
||||||
|
_FUSED_STEP_CACHE["step"] = step
|
||||||
|
return step
|
||||||
|
|
||||||
|
|
||||||
def integrate_batch_discharge(
|
def integrate_batch_discharge(
|
||||||
xp,
|
xp,
|
||||||
q0,
|
q0,
|
||||||
@@ -139,6 +191,11 @@ def integrate_batch_discharge(
|
|||||||
energy_diss = xp.zeros(n, dtype=xp.float64)
|
energy_diss = xp.zeros(n, dtype=xp.float64)
|
||||||
|
|
||||||
_old_err = xp.seterr(all="ignore") if hasattr(xp, "seterr") else None # стиффные конфиги переполняют fixed-step
|
_old_err = xp.seterr(all="ignore") if hasattr(xp, "seterr") else None # стиффные конфиги переполняют fixed-step
|
||||||
|
# на cupy — слитое в одно ядро RK4-ядро (cupy.fuse), на numpy — обычный путь
|
||||||
|
is_cupy = xp.__name__ == "cupy"
|
||||||
|
if is_cupy:
|
||||||
|
fused = _get_fused_step(xp)
|
||||||
|
hs = (params.coil_length_m + params.slug_length_m) / 2
|
||||||
# проверку «все ли готовы» делаем НЕ каждый шаг: на GPU это device->host
|
# проверку «все ли готовы» делаем НЕ каждый шаг: на GPU это device->host
|
||||||
# синхронизация, которая убивает конвейер. Раз в SYNC_EVERY шагов достаточно.
|
# синхронизация, которая убивает конвейер. Раз в SYNC_EVERY шагов достаточно.
|
||||||
SYNC_EVERY = 256
|
SYNC_EVERY = 256
|
||||||
@@ -149,20 +206,25 @@ def integrate_batch_discharge(
|
|||||||
|
|
||||||
q_old, i_old, x_old, v_old = q, i, x, v
|
q_old, i_old, x_old, v_old = q, i, x, v
|
||||||
|
|
||||||
# RK4 от старого состояния
|
if is_cupy:
|
||||||
k1 = _derivatives(xp, q_old, i_old, x_old, v_old, params)
|
q_new, i_new, x_new, v_new, overlap_old, overlap_new = fused(
|
||||||
k2 = _derivatives(xp, q_old + dt / 2 * k1[0], i_old + dt / 2 * k1[1], x_old + dt / 2 * k1[2], v_old + dt / 2 * k1[3], params)
|
q_old, i_old, x_old, v_old, dt, hs, params.smoothing_width_m,
|
||||||
k3 = _derivatives(xp, q_old + dt / 2 * k2[0], i_old + dt / 2 * k2[1], x_old + dt / 2 * k2[2], v_old + dt / 2 * k2[3], params)
|
params.i_sat_a, params.l_air_h, params.l_iron_coeff, params.capacitance_f,
|
||||||
k4 = _derivatives(xp, q_old + dt * k3[0], i_old + dt * k3[1], x_old + dt * k3[2], v_old + dt * k3[3], params)
|
params.r_total_ohm, params.r_eddy_coeff_ohm, params.mass_kg,
|
||||||
|
)
|
||||||
q_new = q_old + dt / 6 * (k1[0] + 2 * k2[0] + 2 * k3[0] + k4[0])
|
else:
|
||||||
i_new = i_old + dt / 6 * (k1[1] + 2 * k2[1] + 2 * k3[1] + k4[1])
|
k1 = _derivatives(xp, q_old, i_old, x_old, v_old, params)
|
||||||
x_new = x_old + dt / 6 * (k1[2] + 2 * k2[2] + 2 * k3[2] + k4[2])
|
k2 = _derivatives(xp, q_old + dt / 2 * k1[0], i_old + dt / 2 * k1[1], x_old + dt / 2 * k1[2], v_old + dt / 2 * k1[3], params)
|
||||||
v_new = v_old + dt / 6 * (k1[3] + 2 * k2[3] + 2 * k3[3] + k4[3])
|
k3 = _derivatives(xp, q_old + dt / 2 * k2[0], i_old + dt / 2 * k2[1], x_old + dt / 2 * k2[2], v_old + dt / 2 * k2[3], params)
|
||||||
|
k4 = _derivatives(xp, q_old + dt * k3[0], i_old + dt * k3[1], x_old + dt * k3[2], v_old + dt * k3[3], params)
|
||||||
|
q_new = q_old + dt / 6 * (k1[0] + 2 * k2[0] + 2 * k3[0] + k4[0])
|
||||||
|
i_new = i_old + dt / 6 * (k1[1] + 2 * k2[1] + 2 * k3[1] + k4[1])
|
||||||
|
x_new = x_old + dt / 6 * (k1[2] + 2 * k2[2] + 2 * k3[2] + k4[2])
|
||||||
|
v_new = v_old + dt / 6 * (k1[3] + 2 * k2[3] + 2 * k3[3] + k4[3])
|
||||||
|
overlap_old, _ = _overlap(xp, x_old, params)
|
||||||
|
overlap_new, _ = _overlap(xp, x_new, params)
|
||||||
|
|
||||||
# потери I²·R_eff (с вихревыми × overlap) и пиковый ток — только для активных
|
# потери I²·R_eff (с вихревыми × overlap) и пиковый ток — только для активных
|
||||||
overlap_old, _ = _overlap(xp, x_old, params)
|
|
||||||
overlap_new, _ = _overlap(xp, x_new, params)
|
|
||||||
r_eff_old = params.r_total_ohm + params.r_eddy_coeff_ohm * overlap_old
|
r_eff_old = params.r_total_ohm + params.r_eddy_coeff_ohm * overlap_old
|
||||||
r_eff_new = params.r_total_ohm + params.r_eddy_coeff_ohm * overlap_new
|
r_eff_new = params.r_total_ohm + params.r_eddy_coeff_ohm * overlap_new
|
||||||
step_diss = 0.5 * (i_old**2 * r_eff_old + i_new**2 * r_eff_new) * dt
|
step_diss = 0.5 * (i_old**2 * r_eff_old + i_new**2 * r_eff_new) * dt
|
||||||
|
|||||||
Reference in New Issue
Block a user