GPU: вся cut-логика шага слита в одно cupy.fuse-ядро

Профиль на 1070: RK4 был слит, но пошаговая бухгалтерия (_magnetic_energy,
_overlap, ~15 where) запускала десятки мелких ядер и съедала >90% времени
(10 геномов/с). Теперь весь шаг — одно ядро; при несборке fuse — честный
fallback на прежний путь. numpy-путь не тронут (эталон для тестов).

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
jze9
2026-07-08 03:42:05 +05:00
parent 1480f4ce1a
commit 965c18f327

View File

@@ -167,6 +167,97 @@ def _get_fused_step(cp):
return step
def _get_fused_full_step(cp):
"""RK4-шаг + ВСЯ пошаговая cut-логика (обрывы, потери, пики) одним ядром.
Профиль на GTX 1070 показал: сам RK4 слит, но бухгалтерия обрыва
(_magnetic_energy, _overlap, ~15 where на шаг) запускала десятки мелких
ядер и съедала >90% времени. Здесь всё элементно и фьюзится в одно ядро.
"""
if "full" in _FUSED_STEP_CACHE:
return _FUSED_STEP_CACHE["full"]
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
def _der(q, i, x, v, hs, w, isat, lair, liron, C, rt, red, fc, dc, m):
s1 = 0.5 * (1.0 + cp.tanh((x + hs) / w * 0.5))
s2 = 0.5 * (1.0 + cp.tanh((hs - x) / w * 0.5))
ov = s1 * s2
dov = (s1 * s2 / w) * (s2 - s1)
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 - cp.tanh(v * 100.0) * (fc + dc * v * v)) / m
return dq, di, v, dv
@cp.fuse()
def full_step(
q, i, x, v, done, committed, past_peak, exit_v, exit_x, exit_q, peak_i, e_diss,
dt, hs, w, isat, lair, liron, C, rt, red, fc, dc, m,
):
a = _der(q, i, x, v, hs, w, isat, lair, liron, C, rt, red, fc, dc, 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, fc, dc, 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, fc, dc, 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, fc, dc, 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])
active = ~done
ov_old = _ovl(x, hs, w)
ov_new = _ovl(xn, hs, w)
# потери шага: I²R_eff + работа трения/воздуха (трапеция)
step_diss = 0.5 * (i * i * (rt + red * ov_old) + in_ * in_ * (rt + red * ov_new)) * dt
fric = 0.5 * ((fc + dc * v * v) * cp.abs(v) + (fc + dc * vn * vn) * cp.abs(vn)) * dt
zero = 0.0 * q
e_diss2 = e_diss + cp.where(active, step_diss + fric, zero)
peak2 = cp.maximum(peak_i, cp.where(active, cp.abs(in_), zero))
past2 = past_peak | (active & (i > 1.0) & (in_ < i))
crossed = active & (i > 0) & (in_ <= 0)
decayed = active & past2 & (in_ < 0.1) & ~crossed # ток < удержания ключа
local_min = active & past2 & (in_ > i) & ~crossed & ~decayed
cut_state = local_min | decayed # обрыв в состоянии ДО шага
cut = crossed | cut_state
frac = cp.where(crossed, i / (i - in_ + 1e-30), zero)
exit_v2 = cp.where(crossed, v + frac * (vn - v), cp.where(cut_state, v, exit_v))
exit_x2 = cp.where(crossed, x + frac * (xn - x), cp.where(cut_state, x, exit_x))
exit_q2 = cp.where(crossed, q + frac * (qn - q), cp.where(cut_state, q, exit_q))
# остаточная магнитная энергия (с насыщением) при обрыве -> freewheel
z = i / isat
g = isat * cp.tanh(z)
az = cp.abs(z)
G = isat * isat * (az + cp.log1p(cp.exp(-2.0 * az)) - LN2)
w_mag = 0.5 * lair * i * i + liron * ov_old * (i * g - G)
e_diss3 = e_diss2 + cp.where(cut_state, w_mag, zero)
committed2 = committed | cut
blew = active & ~cut & ~(cp.abs(in_) < 1e30) # inf/nan: стиффный взрыв
done2 = done | cut | blew
adv = active & ~blew
q2 = cp.where(adv, qn, q)
i2 = cp.where(adv, in_, i)
x2 = cp.where(adv, xn, x)
v2 = cp.where(adv, vn, v)
return q2, i2, x2, v2, done2, committed2, past2, exit_v2, exit_x2, exit_q2, peak2, e_diss3
_FUSED_STEP_CACHE["full"] = full_step
return full_step
def integrate_batch_discharge(
xp,
q0,
@@ -197,9 +288,14 @@ def integrate_batch_discharge(
energy_diss = xp.zeros(n, dtype=xp.float64)
_old_err = xp.seterr(all="ignore") if hasattr(xp, "seterr") else None # стиффные конфиги переполняют fixed-step
# на cupy — слитое в одно ядро RK4-ядро (cupy.fuse), на numpy — обычный путь
# на cupy — ВЕСЬ шаг (RK4 + cut-логика) одним fused-ядром; numpy — обычный путь
is_cupy = xp.__name__ == "cupy"
fused_full = None
if is_cupy:
try:
fused_full = _get_fused_full_step(xp)
except Exception:
fused_full = None # честный fallback на пошаговый путь
fused = _get_fused_step(xp)
hs = (params.coil_length_m + params.slug_length_m) / 2
# проверку «все ли готовы» делаем НЕ каждый шаг: на GPU это device->host
@@ -212,6 +308,25 @@ def integrate_batch_discharge(
q_old, i_old, x_old, v_old = q, i, x, v
if fused_full is not None:
# весь шаг (RK4 + бухгалтерия обрывов/потерь) — ОДНО ядро
try:
(q, i, x, v, done, committed, past_peak,
exit_v, exit_x, exit_q, peak_current, energy_diss) = fused_full(
q_old, i_old, x_old, v_old, done, committed, past_peak,
exit_v, exit_x, exit_q, peak_current, energy_diss,
dt, hs, params.smoothing_width_m, params.i_sat_a,
params.l_air_h, params.l_iron_coeff, params.capacitance_f,
params.r_total_ohm, params.r_eddy_coeff_ohm,
params.retard_const_n, params.drag_coeff_n, params.mass_kg,
)
continue
except Exception:
if step == 0: # fuse не собрался — честный откат на пошаговый путь
fused_full = None
else:
raise
if is_cupy:
q_new, i_new, x_new, v_new, overlap_old, overlap_new = fused(
q_old, i_old, x_old, v_old, dt, hs, params.smoothing_width_m,