GPU: cut-логика двумя fused-ядрами (лимит параметров CUDA 4096 байт)
Одно ядро на весь шаг переполняло formal parameter space (4272>4096). Разбито на step_phys (RK4+потери+ост.энергия) и step_cut (обрывы, без физпараметров): 2 запуска ядра на шаг вместо ~20. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -167,22 +167,20 @@ def _get_fused_step(cp):
|
|||||||
return step
|
return step
|
||||||
|
|
||||||
|
|
||||||
def _get_fused_full_step(cp):
|
def _get_fused_pair(cp):
|
||||||
"""RK4-шаг + ВСЯ пошаговая cut-логика (обрывы, потери, пики) одним ядром.
|
"""Два слитых ядра на шаг: RK4+физика и cut-логика.
|
||||||
|
|
||||||
Профиль на GTX 1070 показал: сам RK4 слит, но бухгалтерия обрыва
|
Профиль на 1070: RK4 был слит, но бухгалтерия обрыва (_magnetic_energy,
|
||||||
(_magnetic_energy, _overlap, ~15 where на шаг) запускала десятки мелких
|
_overlap, ~15 where на шаг) запускала десятки мелких ядер и съедала >90%
|
||||||
ядер и съедала >90% времени. Здесь всё элементно и фьюзится в одно ядро.
|
времени. Одним ядром не влезает в лимит параметров CUDA (4096 байт,
|
||||||
|
«Formal parameter space overflowed»), поэтому два сбалансированных:
|
||||||
|
step_phys — RK4 + потери шага + остаточная магнитная энергия;
|
||||||
|
step_cut — обрывы/пики/выходные состояния, вообще без физических параметров.
|
||||||
"""
|
"""
|
||||||
if "full" in _FUSED_STEP_CACHE:
|
if "pair" in _FUSED_STEP_CACHE:
|
||||||
return _FUSED_STEP_CACHE["full"]
|
return _FUSED_STEP_CACHE["pair"]
|
||||||
LN2 = 0.6931471805599453
|
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):
|
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))
|
s1 = 0.5 * (1.0 + cp.tanh((x + hs) / w * 0.5))
|
||||||
s2 = 0.5 * (1.0 + cp.tanh((hs - x) / w * 0.5))
|
s2 = 0.5 * (1.0 + cp.tanh((hs - x) / w * 0.5))
|
||||||
@@ -202,10 +200,7 @@ def _get_fused_full_step(cp):
|
|||||||
return dq, di, v, dv
|
return dq, di, v, dv
|
||||||
|
|
||||||
@cp.fuse()
|
@cp.fuse()
|
||||||
def full_step(
|
def step_phys(q, i, x, v, dt, hs, w, isat, lair, liron, C, rt, red, fc, dc, m):
|
||||||
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)
|
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)
|
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)
|
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)
|
||||||
@@ -214,36 +209,41 @@ def _get_fused_full_step(cp):
|
|||||||
in_ = i + dt / 6.0 * (a[1] + 2 * b[1] + 2 * c[1] + d[1])
|
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])
|
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])
|
vn = v + dt / 6.0 * (a[3] + 2 * b[3] + 2 * c[3] + d[3])
|
||||||
|
s1o = 0.5 * (1.0 + cp.tanh((x + hs) / w * 0.5))
|
||||||
active = ~done
|
s2o = 0.5 * (1.0 + cp.tanh((hs - x) / w * 0.5))
|
||||||
ov_old = _ovl(x, hs, w)
|
ov_old = s1o * s2o
|
||||||
ov_new = _ovl(xn, hs, w)
|
s1n = 0.5 * (1.0 + cp.tanh((xn + hs) / w * 0.5))
|
||||||
# потери шага: I²R_eff + работа трения/воздуха (трапеция)
|
s2n = 0.5 * (1.0 + cp.tanh((hs - xn) / w * 0.5))
|
||||||
|
ov_new = s1n * s2n
|
||||||
|
# потери шага: I²R_eff (вихревые × overlap) + трение/воздух (трапеция)
|
||||||
step_diss = 0.5 * (i * i * (rt + red * ov_old) + in_ * in_ * (rt + red * ov_new)) * dt
|
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
|
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
|
z = i / isat
|
||||||
g = isat * cp.tanh(z)
|
g = isat * cp.tanh(z)
|
||||||
az = cp.abs(z)
|
az = cp.abs(z)
|
||||||
G = isat * isat * (az + cp.log1p(cp.exp(-2.0 * az)) - LN2)
|
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)
|
w_mag = 0.5 * lair * i * i + liron * ov_old * (i * g - G)
|
||||||
e_diss3 = e_diss2 + cp.where(cut_state, w_mag, zero)
|
return qn, in_, xn, vn, step_diss + fric, w_mag
|
||||||
|
|
||||||
|
@cp.fuse()
|
||||||
|
def step_cut(q, i, x, v, qn, in_, xn, vn, sdiss, w_mag,
|
||||||
|
done, committed, past_peak, exit_v, exit_x, exit_q, peak_i, e_diss):
|
||||||
|
active = ~done
|
||||||
|
zero = 0.0 * q
|
||||||
|
e2 = e_diss + cp.where(active, sdiss, 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))
|
||||||
|
e3 = e2 + cp.where(cut_state, w_mag, zero) # freewheel-диод
|
||||||
committed2 = committed | cut
|
committed2 = committed | cut
|
||||||
blew = active & ~cut & ~(cp.abs(in_) < 1e30) # inf/nan: стиффный взрыв
|
blew = active & ~cut & ~(cp.abs(in_) < 1e30) # inf/nan: стиффный взрыв
|
||||||
done2 = done | cut | blew
|
done2 = done | cut | blew
|
||||||
@@ -252,10 +252,10 @@ def _get_fused_full_step(cp):
|
|||||||
i2 = cp.where(adv, in_, i)
|
i2 = cp.where(adv, in_, i)
|
||||||
x2 = cp.where(adv, xn, x)
|
x2 = cp.where(adv, xn, x)
|
||||||
v2 = cp.where(adv, vn, v)
|
v2 = cp.where(adv, vn, v)
|
||||||
return q2, i2, x2, v2, done2, committed2, past2, exit_v2, exit_x2, exit_q2, peak2, e_diss3
|
return q2, i2, x2, v2, done2, committed2, past2, exit_v2, exit_x2, exit_q2, peak2, e3
|
||||||
|
|
||||||
_FUSED_STEP_CACHE["full"] = full_step
|
_FUSED_STEP_CACHE["pair"] = (step_phys, step_cut)
|
||||||
return full_step
|
return _FUSED_STEP_CACHE["pair"]
|
||||||
|
|
||||||
|
|
||||||
def integrate_batch_discharge(
|
def integrate_batch_discharge(
|
||||||
@@ -290,12 +290,12 @@ def integrate_batch_discharge(
|
|||||||
_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 + cut-логика) одним fused-ядром; numpy — обычный путь
|
# на cupy — ВЕСЬ шаг (RK4 + cut-логика) одним fused-ядром; numpy — обычный путь
|
||||||
is_cupy = xp.__name__ == "cupy"
|
is_cupy = xp.__name__ == "cupy"
|
||||||
fused_full = None
|
fused_pair = None
|
||||||
if is_cupy:
|
if is_cupy:
|
||||||
try:
|
try:
|
||||||
fused_full = _get_fused_full_step(xp)
|
fused_pair = _get_fused_pair(xp)
|
||||||
except Exception:
|
except Exception:
|
||||||
fused_full = None # честный fallback на пошаговый путь
|
fused_pair = None # честный fallback на пошаговый путь
|
||||||
fused = _get_fused_step(xp)
|
fused = _get_fused_step(xp)
|
||||||
hs = (params.coil_length_m + params.slug_length_m) / 2
|
hs = (params.coil_length_m + params.slug_length_m) / 2
|
||||||
# проверку «все ли готовы» делаем НЕ каждый шаг: на GPU это device->host
|
# проверку «все ли готовы» делаем НЕ каждый шаг: на GPU это device->host
|
||||||
@@ -308,22 +308,24 @@ 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
|
||||||
|
|
||||||
if fused_full is not None:
|
if fused_pair is not None:
|
||||||
# весь шаг (RK4 + бухгалтерия обрывов/потерь) — ОДНО ядро
|
# весь шаг = 2 ядра: RK4+физика, затем cut-логика (вместо ~20 мелких)
|
||||||
try:
|
try:
|
||||||
(q, i, x, v, done, committed, past_peak,
|
qn, in_, xn, vn, sdiss, w_mag = fused_pair[0](
|
||||||
exit_v, exit_x, exit_q, peak_current, energy_diss) = fused_full(
|
q_old, i_old, x_old, v_old, dt, hs, params.smoothing_width_m,
|
||||||
q_old, i_old, x_old, v_old, done, committed, past_peak,
|
params.i_sat_a, params.l_air_h, params.l_iron_coeff, params.capacitance_f,
|
||||||
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.r_total_ohm, params.r_eddy_coeff_ohm,
|
||||||
params.retard_const_n, params.drag_coeff_n, params.mass_kg,
|
params.retard_const_n, params.drag_coeff_n, params.mass_kg,
|
||||||
)
|
)
|
||||||
|
(q, i, x, v, done, committed, past_peak,
|
||||||
|
exit_v, exit_x, exit_q, peak_current, energy_diss) = fused_pair[1](
|
||||||
|
q_old, i_old, x_old, v_old, qn, in_, xn, vn, sdiss, w_mag,
|
||||||
|
done, committed, past_peak, exit_v, exit_x, exit_q, peak_current, energy_diss,
|
||||||
|
)
|
||||||
continue
|
continue
|
||||||
except Exception:
|
except Exception:
|
||||||
if step == 0: # fuse не собрался — честный откат на пошаговый путь
|
if step == 0: # fuse не собрался — честный откат на пошаговый путь
|
||||||
fused_full = None
|
fused_pair = None
|
||||||
else:
|
else:
|
||||||
raise
|
raise
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user