Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
86 changes: 47 additions & 39 deletions src/torchsim/sequence/_epg_triton.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,23 +97,16 @@ def _sincos(x):
would each make their own.
"""
quarter = tl.extra.cuda.libdevice.rint(x * 0.6366197723675814)
r = x - quarter * 1.5703125
r = r - quarter * 4.837512969970703125e-4
r = r - quarter * 7.54978995489188216e-8
r = tl.fma(-quarter, 1.5703125, x)
r = tl.fma(-quarter, 4.837512969970703125e-4, r)
r = tl.fma(-quarter, 7.54978995489188216e-8, r)
r2 = r * r
sine = r + r * r2 * (
-1.6666654611e-1 + r2 * (8.3321608736e-3 + r2 * -1.9515295891e-4)
)
cosine = (
1.0
- 0.5 * r2
+ r2
* r2
* (
4.166664568298827e-2
+ r2 * (-1.388731625493765e-3 + r2 * 2.443315711809948e-5)
)
)
sine = tl.fma(-1.9515295891e-4, r2, 8.3321608736e-3)
sine = tl.fma(sine, r2, -1.6666654611e-1)
sine = tl.fma(r * r2, sine, r)
cosine = tl.fma(2.443315711809948e-5, r2, -1.388731625493765e-3)
cosine = tl.fma(cosine, r2, 4.166664568298827e-2)
cosine = tl.fma(r2 * r2, cosine, tl.fma(-0.5, r2, 1.0))
q = quarter.to(tl.int32) & 3
s = tl.where(
q == 0, sine, tl.where(q == 1, cosine, tl.where(q == 2, -sine, -cosine))
Expand Down Expand Up @@ -14679,27 +14672,42 @@ def _rotate_flip_phase(
"""
cosine_half_sq = 0.5 * (1.0 + cosine)
sine_half_sq = 0.5 * (1.0 - cosine)
half_sine = 0.5 * sine

# Every sum of products is one fused multiply-add in a fixed order, so a
# kernel compiled for fewer terms rounds the rotation as the full one does.
minus_2phi_r = tl.fma(cos_2phi, fm_r, -(sin_2phi * fm_i))
minus_2phi_i = tl.fma(sin_2phi, fm_r, cos_2phi * fm_i)
plus_2phi_r = tl.fma(cos_2phi, fp_r, sin_2phi * fp_i)
plus_2phi_i = tl.fma(cos_2phi, fp_i, -(sin_2phi * fp_r))
z_turn_a = tl.fma(sin_phi, z_r, cos_phi * z_i)
z_turn_b = tl.fma(sin_phi, z_i, -(cos_phi * z_r))
z_turn_c = tl.fma(sin_phi, z_r, -(cos_phi * z_i))
z_turn_d = tl.fma(cos_phi, z_r, sin_phi * z_i)

rotated_pr = tl.fma(
sine, z_turn_a, tl.fma(sine_half_sq, minus_2phi_r, cosine_half_sq * fp_r)
)
rotated_pi = tl.fma(
sine, z_turn_b, tl.fma(sine_half_sq, minus_2phi_i, cosine_half_sq * fp_i)
)
rotated_mr = tl.fma(
sine, z_turn_c, tl.fma(cosine_half_sq, fm_r, sine_half_sq * plus_2phi_r)
)
rotated_mi = tl.fma(
sine, z_turn_d, tl.fma(cosine_half_sq, fm_i, sine_half_sq * plus_2phi_i)
)

rotated_pr = cosine_half_sq * fp_r
rotated_pr += sine_half_sq * (cos_2phi * fm_r - sin_2phi * fm_i)
rotated_pr += sine * (sin_phi * z_r + cos_phi * z_i)
rotated_pi = cosine_half_sq * fp_i
rotated_pi += sine_half_sq * (sin_2phi * fm_r + cos_2phi * fm_i)
rotated_pi += sine * (sin_phi * z_i - cos_phi * z_r)

rotated_mr = sine_half_sq * (cos_2phi * fp_r + sin_2phi * fp_i)
rotated_mr += cosine_half_sq * fm_r
rotated_mr += sine * (sin_phi * z_r - cos_phi * z_i)
rotated_mi = sine_half_sq * (-sin_2phi * fp_r + cos_2phi * fp_i)
rotated_mi += cosine_half_sq * fm_i
rotated_mi += sine * (cos_phi * z_r + sin_phi * z_i)

rotated_zr = -0.5 * sine * (sin_phi * fp_r - cos_phi * fp_i)
rotated_zr += -0.5 * sine * (sin_phi * fm_r + cos_phi * fm_i)
rotated_zr += cosine * z_r
rotated_zi = -0.5 * sine * (cos_phi * fp_r + sin_phi * fp_i)
rotated_zi += 0.5 * sine * (cos_phi * fm_r - sin_phi * fm_i)
rotated_zi += cosine * z_i
plus_turn_r = tl.fma(sin_phi, fp_r, -(cos_phi * fp_i))
minus_turn_r = tl.fma(sin_phi, fm_r, cos_phi * fm_i)
plus_turn_i = tl.fma(cos_phi, fp_r, sin_phi * fp_i)
minus_turn_i = tl.fma(cos_phi, fm_r, -(sin_phi * fm_i))
rotated_zr = tl.fma(
cosine, z_r, tl.fma(-half_sine, minus_turn_r, -half_sine * plus_turn_r)
)
rotated_zi = tl.fma(
cosine, z_i, tl.fma(half_sine, minus_turn_i, -half_sine * plus_turn_i)
)
return (
rotated_pr,
rotated_pi,
Expand Down Expand Up @@ -15201,8 +15209,8 @@ def _epg_kernel(
sin_event = _event_value(
phase_sin, event_base, event, active_atom, single_train
)
cos_phi = cos_event * b1_cos - sin_event * b1_sin
sin_phi = sin_event * b1_cos + cos_event * b1_sin
cos_phi = tl.fma(cos_event, b1_cos, -(sin_event * b1_sin))
sin_phi = tl.fma(sin_event, b1_cos, cos_event * b1_sin)
if profiled or dynamic:
# Either pair is built at zero RF phase, which turns the rotation
# axis and so reaches ``b`` alone.
Expand Down Expand Up @@ -15245,7 +15253,7 @@ def _epg_kernel(
)
)
sine, cosine = _sincos(alpha)
cos_2phi = cos_phi * cos_phi - sin_phi * sin_phi
cos_2phi = tl.fma(cos_phi, cos_phi, -(sin_phi * sin_phi))
sin_2phi = 2.0 * sin_phi * cos_phi

(
Expand Down
Loading