diff --git a/src/torchsim/sequence/_epg_triton.py b/src/torchsim/sequence/_epg_triton.py index 80d09f24..1bf67e0b 100644 --- a/src/torchsim/sequence/_epg_triton.py +++ b/src/torchsim/sequence/_epg_triton.py @@ -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)) @@ -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, @@ -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. @@ -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 (