diff --git a/src/torchsim/sequence/_epg_triton.py b/src/torchsim/sequence/_epg_triton.py index ebb00780..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. @@ -14698,6 +14735,8 @@ def _epg_kernel( kind, flip, phase, + phase_cos, + phase_sin, action, output_index, shim_index, @@ -14839,6 +14878,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( @@ -15093,35 +15134,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 +15186,75 @@ 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, + 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 + ) + # 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. + 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 = 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) = ( + _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) + sine, cosine = _sincos(alpha) + cos_2phi = cos_phi * cos_phi - sin_phi * sin_phi + sin_2phi = 2.0 * sin_phi * cos_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, - ) - - 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 +15262,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 +15278,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 +15292,119 @@ 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_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 + ) + # 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 @@ -17258,6 +17295,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, @@ -17311,6 +17353,8 @@ def simulate_into( kind, flip, phase, + phase_cos, + phase_sin, action, output_index, shim_index, 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