From 57392e58ff71641fd6f7bc5c355670de05c23b2a Mon Sep 17 00:00:00 2001 From: mcencini Date: Tue, 6 Oct 2026 16:44:29 +0200 Subject: [PATCH] Round a pulse the same whichever terms the kernel was compiled for The pulse's sine and cosine, its phase turned by the transmit field's, its double angle and the rotation of the states are written as fused multiply-adds in a fixed order. Left to the compiler, each kernel variant contracted them its own way, and a tissue declaring fewer terms got an answer one ulp away from the full kernel's on the same input. Co-Authored-By: Claude Opus 5.5 --- src/torchsim/sequence/_epg_triton.py | 86 +++++++++++++++------------- 1 file changed, 47 insertions(+), 39 deletions(-) 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 (