From c44e34ca05bed346032411a5725b2f14159a69ed Mon Sep 17 00:00:00 2001 From: mcencini Date: Mon, 5 Oct 2026 23:00:24 +0200 Subject: [PATCH 1/4] Lay out off-resonance and transmit phase per voxel together on a card The CUDA kernels carry the two through one launch flag and read both per voxel when either is declared, while the tissue buffers narrowed each to a single value by its own feature. A tissue giving off-resonance beside a transmit map, without a transmit phase, then read the phase past the end of its buffer: a wrong echo phase that changed with the batch. Co-Authored-By: Claude Opus 5.5 --- src/torchsim/sequence/_simulation.py | 9 ++++ tests/sequence/test_cuda_parity.py | 67 ++++++++++++++++++++++++++++ 2 files changed, 76 insertions(+) diff --git a/src/torchsim/sequence/_simulation.py b/src/torchsim/sequence/_simulation.py index 2685cf31..730edb80 100644 --- a/src/torchsim/sequence/_simulation.py +++ b/src/torchsim/sequence/_simulation.py @@ -65,6 +65,11 @@ # relaxation times are the floor of the model rather than a term it switches on. _ALWAYS_PER_VOXEL = frozenset(("t1_ms", "t2_ms")) +# Features one launch flag carries together, so that a kernel carrying either +# reads every property of both: off-resonance and transmit phase are the one +# ``off_axis`` turn (see ``feature_flags``). +_READ_TOGETHER = (frozenset(("B0", "B1_PHASE")),) + # Which fraction gates each exchanging pool, pool B's first. _POOL_FRACTIONS = ( "pool_b_fraction", @@ -686,6 +691,10 @@ def _prepare_tissue( if parameter.name not in _ALWAYS_PER_VOXEL and not (shims > 1 and index in _TRANSMIT) ) + if features is not None: + for together in _READ_TOGETHER: + if together & features: + features = features | together unread = frozenset( index for index in spare diff --git a/tests/sequence/test_cuda_parity.py b/tests/sequence/test_cuda_parity.py index ca1d6adf..e734c2c6 100644 --- a/tests/sequence/test_cuda_parity.py +++ b/tests/sequence/test_cuda_parity.py @@ -559,3 +559,70 @@ def test_repeated_second_order_runs_agree_to_tolerance(): second = _second_order("cuda", 17, 5) assert _worst_disagreement(first, second) < 1e-5 + + +@pytest.mark.parametrize( + "given", + [("b0_hz",), ("b0_hz", "b1"), ("b1_phase_rad",), ("b1_phase_rad", "b1")], +) +def test_one_half_of_the_off_axis_turn_matches_the_cpu_kernel(given): + """A tissue giving off-resonance or transmit phase alone, beside a map. + + The two share one launch flag, so a kernel carrying either reads both per + voxel; a run that declares one is given room for the other. The spoiler + after each echo winds unlike the crushers, which keeps off-resonance in + the states rather than on the samples. + """ + from torchsim.sequence import ( + AdcRole, + EpgEngine, + EventAction, + EventType, + RfUse, + SequenceDescription, + SequenceEvent, + ideal_rf_definition, + ) + + events = [] + for repetition in range(3): + start = 500_000.0 * repetition + events += [ + SequenceEvent.rf(start + 1_000.0, 0, RfUse.EXCITATION, torch.pi / 2, 0.0), + SequenceEvent(EventType.WAIT, start + 3_000.0, (), EventAction.SHIFT_AFTER), + SequenceEvent.rf( + start + 7_000.0, 0, RfUse.REFOCUSING, torch.pi, torch.pi / 2 + ), + SequenceEvent(EventType.WAIT, start + 9_000.0, (), EventAction.SHIFT_AFTER), + SequenceEvent.adc(start + 13_000.0, AdcRole.SINGLE, 0.0), + SequenceEvent( + EventType.WAIT, start + 20_000.0, (), EventAction.SHIFT_AFTER + ), + ] + description = SequenceDescription( + 0, 1_500_000.0, tuple(events), {0: ideal_rf_definition()} + ) + values = { + "b0_hz": [0.0, 1.0, 5.0, 13.0], + "b1": [1.0, 1.05, 0.9, 1.0], + "b1_phase_rad": [0.0, 0.1, 0.2, 0.3], + } + # Blocks the allocator hands out again, left holding something other than + # zero: a buffer read past its end then reads that, not a lucky zero. + dirty = [torch.full((64,), 7.0, device="cuda") for _ in range(256)] + del dirty + signals = [] + for device in ("cpu", "cuda"): + tissue = TissueProperties( + t1_ms=torch.full((4,), 300.0, device=device), + t2_ms=torch.full((4,), 40.0, device=device), + **{name: torch.tensor(values[name], device=device) for name in given}, + ) + signals.append( + EpgEngine() + .simulate(description, tissue, record="all", device=device) + .signal + ) + expected, actual = signals + scale = expected.abs().max() + assert ((expected - actual.cpu()).abs().max() / scale) < 1e-5 From d3263daa0acde40264a1ac066c5a09cfd03203de Mon Sep 17 00:00:00 2001 From: mcencini Date: Tue, 6 Oct 2026 12:48:12 +0200 Subject: [PATCH 2/4] Let each event of the fused EPG kernel pay only for what it does Every program reads the same event, so a branch on its kind and action is taken uniformly: a pulse that turns rotates the states, a recorded readout reads them, a shift shifts them, and no event computes the others' work to throw it away. On a 7132-event spoiled stream over 6210 trains and atoms the kernel takes 138 ms where it took 193 ms. Co-Authored-By: Claude Opus 5.5 --- src/torchsim/sequence/_epg_triton.py | 408 +++++++++++++-------------- 1 file changed, 199 insertions(+), 209 deletions(-) diff --git a/src/torchsim/sequence/_epg_triton.py b/src/torchsim/sequence/_epg_triton.py index ebb00780..6b3d4536 100644 --- a/src/torchsim/sequence/_epg_triton.py +++ b/src/torchsim/sequence/_epg_triton.py @@ -15093,35 +15093,29 @@ def _epg_kernel( ) longitudinal_real += tl.where(state == 0, recovery, 0.0) + # Every program reads the same event, so a branch on what it does is + # taken by all of them alike: an event pays only for what it does. event_action = tl.load(action + event).to(tl.int32) - pre_shift = (event_action & 1) != 0 - shifted_pr, shifted_pi, shifted_mr, shifted_mi = _shift( - fplus_real, - fplus_imag, - fminus_real, - fminus_imag, - state, - state_mask, - state_count, - ) - fplus_real = tl.where(pre_shift, shifted_pr, fplus_real) - fplus_imag = tl.where(pre_shift, shifted_pi, fplus_imag) - fminus_real = tl.where(pre_shift, shifted_mr, fminus_real) - fminus_imag = tl.where(pre_shift, shifted_mi, fminus_imag) - if pools == 2 or pools == 3: - b_pr, b_pi, b_mr, b_mi = _shift( - bplus_real, - bplus_imag, - bminus_real, - bminus_imag, + if (event_action & 1) != 0: + fplus_real, fplus_imag, fminus_real, fminus_imag = _shift( + fplus_real, + fplus_imag, + fminus_real, + fminus_imag, state, state_mask, state_count, ) - bplus_real = tl.where(pre_shift, b_pr, bplus_real) - bplus_imag = tl.where(pre_shift, b_pi, bplus_imag) - bminus_real = tl.where(pre_shift, b_mr, bminus_real) - bminus_imag = tl.where(pre_shift, b_mi, bminus_imag) + if pools == 2 or pools == 3: + bplus_real, bplus_imag, bminus_real, bminus_imag = _shift( + bplus_real, + bplus_imag, + bminus_real, + bminus_imag, + state, + state_mask, + state_count, + ) event_kind = tl.load(kind + event) is_rf = event_kind == 1 @@ -15151,91 +15145,70 @@ def _epg_kernel( atom_b1_phase = tl.load( b1_phase + row + atom, mask=active_atom, other=0.0 ) - alpha = ( - _event_value(flip, event_base, event, active_atom, single_train) * atom_b1 - ) - phi = ( - _event_value(phase, event_base, event, active_atom, single_train) - + atom_b1_phase - ) - if profiled or dynamic: - # Either pair is built at zero RF phase, which turns the rotation - # axis and so reaches ``b`` alone. - if dynamic: - # Already integrated at this pulse's own flip, so the flip is - # inside the pair rather than read against it. - pair = _dynamic_pair_at( - pairs, - pair_index, - event_base, - event, - atom, - atom_count, - active_atom, - ) - else: - pair = _profile_pair( - profile, - _table_row(profile_index, event, location, locations), - alpha, - profile_bins, - profile_step, - ) - turn_r = tl.cos(phi) - turn_i = -tl.sin(phi) - spun_br = pair[2] * turn_r - pair[3] * turn_i - spun_bi = pair[2] * turn_i + pair[3] * turn_r - (shaped_pr, shaped_pi, shaped_mr, shaped_mi, shaped_zr, shaped_zi) = ( - _rotate_spinor( - pair[0], - pair[1], - spun_br, - spun_bi, - fplus_real, - fplus_imag, - fminus_real, - fminus_imag, - longitudinal_real, - longitudinal_imag, - ) + if (event_kind == 1) & ((event_action & 4) == 0): + alpha = ( + _event_value(flip, event_base, event, active_atom, single_train) + * atom_b1 ) - cosine = tl.cos(alpha) - sine = tl.sin(alpha) - cos_phi = tl.cos(phi) - sin_phi = tl.sin(phi) - cos_2phi = tl.cos(2.0 * phi) - sin_2phi = tl.sin(2.0 * phi) - - ( - rotated_pr, - rotated_pi, - rotated_mr, - rotated_mi, - rotated_zr, - rotated_zi, - ) = _rotate_flip_phase( - cosine, - sine, - cos_phi, - sin_phi, - cos_2phi, - sin_2phi, - fplus_real, - fplus_imag, - fminus_real, - fminus_imag, - longitudinal_real, - longitudinal_imag, - ) + phi = ( + _event_value(phase, event_base, event, active_atom, single_train) + + atom_b1_phase + ) + if profiled or dynamic: + # Either pair is built at zero RF phase, which turns the rotation + # axis and so reaches ``b`` alone. + if dynamic: + # Already integrated at this pulse's own flip, so the flip is + # inside the pair rather than read against it. + pair = _dynamic_pair_at( + pairs, + pair_index, + event_base, + event, + atom, + atom_count, + active_atom, + ) + else: + pair = _profile_pair( + profile, + _table_row(profile_index, event, location, locations), + alpha, + profile_bins, + profile_step, + ) + turn_r = tl.cos(phi) + turn_i = -tl.sin(phi) + spun_br = pair[2] * turn_r - pair[3] * turn_i + spun_bi = pair[2] * turn_i + pair[3] * turn_r + (shaped_pr, shaped_pi, shaped_mr, shaped_mi, shaped_zr, shaped_zi) = ( + _rotate_spinor( + pair[0], + pair[1], + spun_br, + spun_bi, + fplus_real, + fplus_imag, + fminus_real, + fminus_imag, + longitudinal_real, + longitudinal_imag, + ) + ) + cosine = tl.cos(alpha) + sine = tl.sin(alpha) + cos_phi = tl.cos(phi) + sin_phi = tl.sin(phi) + cos_2phi = tl.cos(2.0 * phi) + sin_2phi = tl.sin(2.0 * phi) - if pools == 2 or pools == 3: ( - b_rot_pr, - b_rot_pi, - b_rot_mr, - b_rot_mi, - b_rot_zr, - b_rot_zi, + rotated_pr, + rotated_pi, + rotated_mr, + rotated_mi, + rotated_zr, + rotated_zi, ) = _rotate_flip_phase( cosine, sine, @@ -15243,16 +15216,15 @@ def _epg_kernel( sin_phi, cos_2phi, sin_2phi, - bplus_real, - bplus_imag, - bminus_real, - bminus_imag, - bound_real, - bound_imag, + fplus_real, + fplus_imag, + fminus_real, + fminus_imag, + longitudinal_real, + longitudinal_imag, ) - if profiled or dynamic: - # The same pulse, the same rotation: a chemical shift moves - # where a pool precesses, not what a pulse does to it. + + if pools == 2 or pools == 3: ( b_rot_pr, b_rot_pi, @@ -15260,11 +15232,13 @@ def _epg_kernel( b_rot_mi, b_rot_zr, b_rot_zi, - ) = _rotate_spinor( - pair[0], - pair[1], - spun_br, - spun_bi, + ) = _rotate_flip_phase( + cosine, + sine, + cos_phi, + sin_phi, + cos_2phi, + sin_2phi, bplus_real, bplus_imag, bminus_real, @@ -15272,102 +15246,118 @@ def _epg_kernel( bound_real, bound_imag, ) - if profiled or dynamic: - rotated_pr = shaped_pr - rotated_pi = shaped_pi - rotated_mr = shaped_mr - rotated_mi = shaped_mi - rotated_zr = shaped_zr - rotated_zi = shaped_zi + if profiled or dynamic: + # The same pulse, the same rotation: a chemical shift moves + # where a pool precesses, not what a pulse does to it. + ( + b_rot_pr, + b_rot_pi, + b_rot_mr, + b_rot_mi, + b_rot_zr, + b_rot_zi, + ) = _rotate_spinor( + pair[0], + pair[1], + spun_br, + spun_bi, + bplus_real, + bplus_imag, + bminus_real, + bminus_imag, + bound_real, + bound_imag, + ) + if profiled or dynamic: + rotated_pr = shaped_pr + rotated_pi = shaped_pi + rotated_mr = shaped_mr + rotated_mi = shaped_mi + rotated_zr = shaped_zr + rotated_zi = shaped_zi - rotate = is_rf & ~is_inversion - if pools == 2 or pools == 3: - bplus_real = tl.where(rotate, b_rot_pr, bplus_real) - bplus_imag = tl.where(rotate, b_rot_pi, bplus_imag) - bminus_real = tl.where(rotate, b_rot_mr, bminus_real) - bminus_imag = tl.where(rotate, b_rot_mi, bminus_imag) - bound_real = tl.where(rotate, b_rot_zr, bound_real) - bound_imag = tl.where(rotate, b_rot_zi, bound_imag) - if pools == 1 or pools == 3: - # The semisolid pool absorbs the power the pulse deposits, so it - # reads the bare flip the transmit field gives the voxel -- not the - # slice-shaped rotation the free pool takes from the table. - offset = tl.load(rf_frequency + event) - atom_b0 - absorbed = tl.exp( - tl.load(saturation + event) - * alpha - * alpha - * _lineshape_at(lineshape, offset, lineshape_bins, lineshape_step) - ) - if pools == 1: - bound_real = tl.where(rotate, absorbed * bound_real, bound_real) - bound_imag = tl.where(rotate, absorbed * bound_imag, bound_imag) - else: - semisolid_real = tl.where( - rotate, absorbed * semisolid_real, semisolid_real - ) - semisolid_imag = tl.where( - rotate, absorbed * semisolid_imag, semisolid_imag + if pools == 2 or pools == 3: + bplus_real = b_rot_pr + bplus_imag = b_rot_pi + bminus_real = b_rot_mr + bminus_imag = b_rot_mi + bound_real = b_rot_zr + bound_imag = b_rot_zi + if pools == 1 or pools == 3: + # The semisolid pool absorbs the power the pulse deposits, so it + # reads the bare flip the transmit field gives the voxel -- not the + # slice-shaped rotation the free pool takes from the table. + offset = tl.load(rf_frequency + event) - atom_b0 + absorbed = tl.exp( + tl.load(saturation + event) + * alpha + * alpha + * _lineshape_at(lineshape, offset, lineshape_bins, lineshape_step) ) - fplus_real = tl.where(rotate, rotated_pr, fplus_real) - fplus_imag = tl.where(rotate, rotated_pi, fplus_imag) - fminus_real = tl.where(rotate, rotated_mr, fminus_real) - fminus_imag = tl.where(rotate, rotated_mi, fminus_imag) - longitudinal_real = tl.where(rotate, rotated_zr, longitudinal_real) - longitudinal_imag = tl.where(rotate, rotated_zi, longitudinal_imag) - - record = ((event_action & 32) != 0) & (event_kind == 2) - adc_phase = _event_value(phase, event_base, event, active_atom, single_train) - adc_cos = tl.cos(adc_phase) - adc_sin = tl.sin(adc_phase) - # A coil sees the whole voxel, so what it records is the sum over - # pools; each pool's share is already in its own state. - read_real = fplus_real - read_imag = fplus_imag - if pools == 2 or pools == 3: - read_real = fplus_real + bplus_real - read_imag = fplus_imag + bplus_imag - signal_real = atom_m0 * (read_real * adc_cos + read_imag * adc_sin) - signal_imag = atom_m0 * (read_imag * adc_cos - read_real * adc_sin) - out = tl.load(output_index + event) - output_offset = problem * output_count + out - output_mask = active_atom & (state == 0) & record & (out >= 0) - tl.store(output_real + output_offset + state, signal_real, mask=output_mask) - tl.store(output_imag + output_offset + state, signal_imag, mask=output_mask) + if pools == 1: + bound_real = absorbed * bound_real + bound_imag = absorbed * bound_imag + else: + semisolid_real = absorbed * semisolid_real + semisolid_imag = absorbed * semisolid_imag + fplus_real = rotated_pr + fplus_imag = rotated_pi + fminus_real = rotated_mr + fminus_imag = rotated_mi + longitudinal_real = rotated_zr + longitudinal_imag = rotated_zi + + if ((event_action & 32) != 0) & (event_kind == 2): + adc_phase = _event_value( + phase, event_base, event, active_atom, single_train + ) + adc_cos = tl.cos(adc_phase) + adc_sin = tl.sin(adc_phase) + # A coil sees the whole voxel, so what it records is the sum over + # pools; each pool's share is already in its own state. + read_real = fplus_real + read_imag = fplus_imag + if pools == 2 or pools == 3: + read_real = fplus_real + bplus_real + read_imag = fplus_imag + bplus_imag + signal_real = atom_m0 * (read_real * adc_cos + read_imag * adc_sin) + signal_imag = atom_m0 * (read_imag * adc_cos - read_real * adc_sin) + out = tl.load(output_index + event) + output_offset = problem * output_count + out + output_mask = active_atom & (state == 0) & (out >= 0) + tl.store(output_real + output_offset + state, signal_real, mask=output_mask) + tl.store(output_imag + output_offset + state, signal_imag, mask=output_mask) - do_shift = ((event_action & 2) != 0) | ((event_action & 16) != 0) - shifted_pr, shifted_pi, shifted_mr, shifted_mi = _shift( - fplus_real, - fplus_imag, - fminus_real, - fminus_imag, - state, - state_mask, - state_count, - ) - fplus_real = tl.where(do_shift, shifted_pr, fplus_real) - fplus_imag = tl.where(do_shift, shifted_pi, fplus_imag) - fminus_real = tl.where(do_shift, shifted_mr, fminus_real) - fminus_imag = tl.where(do_shift, shifted_mi, fminus_imag) - spoil = (event_action & 8) != 0 - fplus_real = tl.where(spoil, 0.0, fplus_real) - fplus_imag = tl.where(spoil, 0.0, fplus_imag) - fminus_real = tl.where(spoil, 0.0, fminus_real) - fminus_imag = tl.where(spoil, 0.0, fminus_imag) - if pools == 2 or pools == 3: - b_pr, b_pi, b_mr, b_mi = _shift( - bplus_real, - bplus_imag, - bminus_real, - bminus_imag, + if (event_action & 18) != 0: + fplus_real, fplus_imag, fminus_real, fminus_imag = _shift( + fplus_real, + fplus_imag, + fminus_real, + fminus_imag, state, state_mask, state_count, ) - bplus_real = tl.where(spoil, 0.0, tl.where(do_shift, b_pr, bplus_real)) - bplus_imag = tl.where(spoil, 0.0, tl.where(do_shift, b_pi, bplus_imag)) - bminus_real = tl.where(spoil, 0.0, tl.where(do_shift, b_mr, bminus_real)) - bminus_imag = tl.where(spoil, 0.0, tl.where(do_shift, b_mi, bminus_imag)) + if pools == 2 or pools == 3: + bplus_real, bplus_imag, bminus_real, bminus_imag = _shift( + bplus_real, + bplus_imag, + bminus_real, + bminus_imag, + state, + state_mask, + state_count, + ) + if (event_action & 8) != 0: + fplus_real = empty + fplus_imag = empty + fminus_real = empty + fminus_imag = empty + if pools == 2 or pools == 3: + bplus_real = empty + bplus_imag = empty + bminus_real = empty + bminus_imag = empty @triton.jit From 3679375c03c6ab59fc67abb4f30cea96d9a2d0f0 Mon Sep 17 00:00:00 2001 From: mcencini Date: Tue, 6 Oct 2026 12:54:14 +0200 Subject: [PATCH 3/4] Take the cosine and sine of each event's phase once, at the launch Under RF spoiling a phase grows without bound, and every program took the accurate cosine and sine of it at every pulse and every readout. The launch takes them once per event in double precision; a pulse turns them by the transmit field's phase and reads the double angle off their products. The 7132-event stream over 6210 trains and atoms takes 116 ms where it took 138. Co-Authored-By: Claude Opus 5.5 --- src/torchsim/sequence/_epg_triton.py | 44 ++++++++++++++++++++-------- 1 file changed, 31 insertions(+), 13 deletions(-) diff --git a/src/torchsim/sequence/_epg_triton.py b/src/torchsim/sequence/_epg_triton.py index 6b3d4536..4dd0f936 100644 --- a/src/torchsim/sequence/_epg_triton.py +++ b/src/torchsim/sequence/_epg_triton.py @@ -14698,6 +14698,8 @@ def _epg_kernel( kind, flip, phase, + phase_cos, + phase_sin, action, output_index, shim_index, @@ -14839,6 +14841,8 @@ def _epg_kernel( if off_axis: atom_b1_phase = tl.load(b1_phase + scalar_atom, mask=active_atom, other=0.0) atom_b0 = tl.load(b0 + scalar_atom, mask=active_atom, other=0.0) + b1_cos = tl.cos(atom_b1_phase) + b1_sin = tl.sin(atom_b1_phase) atom_inversion = 1.0 if inverting: atom_inversion = tl.load( @@ -15145,15 +15149,23 @@ def _epg_kernel( atom_b1_phase = tl.load( b1_phase + row + atom, mask=active_atom, other=0.0 ) + b1_cos = tl.cos(atom_b1_phase) + b1_sin = tl.sin(atom_b1_phase) if (event_kind == 1) & ((event_action & 4) == 0): alpha = ( _event_value(flip, event_base, event, active_atom, single_train) * atom_b1 ) - phi = ( - _event_value(phase, event_base, event, active_atom, single_train) - + atom_b1_phase + # The pulse's phase, read off the cosine and sine the launch took + # of it, turned by the transmit field's own. + cos_event = _event_value( + phase_cos, event_base, event, active_atom, single_train ) + 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 if profiled or dynamic: # Either pair is built at zero RF phase, which turns the rotation # axis and so reaches ``b`` alone. @@ -15177,8 +15189,8 @@ def _epg_kernel( profile_bins, profile_step, ) - turn_r = tl.cos(phi) - turn_i = -tl.sin(phi) + turn_r = cos_phi + turn_i = -sin_phi spun_br = pair[2] * turn_r - pair[3] * turn_i spun_bi = pair[2] * turn_i + pair[3] * turn_r (shaped_pr, shaped_pi, shaped_mr, shaped_mi, shaped_zr, shaped_zi) = ( @@ -15197,10 +15209,8 @@ def _epg_kernel( ) cosine = tl.cos(alpha) sine = tl.sin(alpha) - cos_phi = tl.cos(phi) - sin_phi = tl.sin(phi) - cos_2phi = tl.cos(2.0 * phi) - sin_2phi = tl.sin(2.0 * phi) + cos_2phi = cos_phi * cos_phi - sin_phi * sin_phi + sin_2phi = 2.0 * sin_phi * cos_phi ( rotated_pr, @@ -15308,11 +15318,12 @@ def _epg_kernel( longitudinal_imag = rotated_zi if ((event_action & 32) != 0) & (event_kind == 2): - adc_phase = _event_value( - phase, event_base, event, active_atom, single_train + adc_cos = _event_value( + phase_cos, event_base, event, active_atom, single_train + ) + adc_sin = _event_value( + phase_sin, event_base, event, active_atom, single_train ) - adc_cos = tl.cos(adc_phase) - adc_sin = tl.sin(adc_phase) # A coil sees the whole voxel, so what it records is the sum over # pools; each pool's share is already in its own state. read_real = fplus_real @@ -17248,6 +17259,11 @@ def simulate_into( tissue, duration, pools=pools, narrow=narrow ) + # Phases grow without bound under RF spoiling, so their cosines and sines + # are taken once here, in double precision, rather than in every program. + phase_cos = torch.cos(phase.double()).to(torch.float32) + phase_sin = torch.sin(phase.double()).to(torch.float32) + if real_axis == 1: _epg_real_kernel[grid]( t1, @@ -17301,6 +17317,8 @@ def simulate_into( kind, flip, phase, + phase_cos, + phase_sin, action, output_index, shim_index, From 46902b9b446f7e70a819ed632aa2bd4f09f3842a Mon Sep 17 00:00:00 2001 From: mcencini Date: Tue, 6 Oct 2026 14:07:17 +0200 Subject: [PATCH 4/4] Take a pulse's flip-angle sine and cosine from one quarter-turn reduction The flip angle times the transmit field is an atom's own, so it cannot be taken at the launch; one Cody-Waite reduction and the single-precision Cephes polynomials give both to about an ulp where two library calls each reduced it. The 6210-train stream takes 99 ms where it took 116. Co-Authored-By: Claude Opus 5.5 --- src/torchsim/sequence/_epg_triton.py | 40 ++++++++++++++++++++++++++-- 1 file changed, 38 insertions(+), 2 deletions(-) diff --git a/src/torchsim/sequence/_epg_triton.py b/src/torchsim/sequence/_epg_triton.py index 4dd0f936..80d09f24 100644 --- a/src/torchsim/sequence/_epg_triton.py +++ b/src/torchsim/sequence/_epg_triton.py @@ -87,6 +87,43 @@ def _first(values, state): return tl.gather(values, index, 1) +@triton.jit +def _sincos(x): + """The sine and cosine of ``x`` from one reduction by a quarter turn. + + Cody and Waite's three-part quarter turn and the single-precision Cephes + polynomials on the eighth turn either side of zero, to about an ulp where + ``|x|`` is a flip angle; one reduction serves both where two library calls + 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 + 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) + ) + ) + q = quarter.to(tl.int32) & 3 + s = tl.where( + q == 0, sine, tl.where(q == 1, cosine, tl.where(q == 2, -sine, -cosine)) + ) + c = tl.where( + q == 0, cosine, tl.where(q == 1, -sine, tl.where(q == 2, -cosine, sine)) + ) + return s, c + + @triton.jit def _event_value(values, event_base, event, active_atom, single_train: tl.constexpr): """One event's entry of a buffer carrying a row per train. @@ -15207,8 +15244,7 @@ def _epg_kernel( longitudinal_imag, ) ) - cosine = tl.cos(alpha) - sine = tl.sin(alpha) + sine, cosine = _sincos(alpha) cos_2phi = cos_phi * cos_phi - sin_phi * sin_phi sin_2phi = 2.0 * sin_phi * cos_phi