diff --git a/src/torchsim/model/_signal.py b/src/torchsim/model/_signal.py index 023014ad..45a275a8 100644 --- a/src/torchsim/model/_signal.py +++ b/src/torchsim/model/_signal.py @@ -36,6 +36,7 @@ __all__ = ["_SignalModel"] +import numbers from abc import ABC, abstractmethod from collections.abc import Mapping, Sequence from copy import copy as shallow_copy @@ -298,8 +299,22 @@ def _shaped(self, signal: torch.Tensor, batch: tuple[int, ...]) -> torch.Tensor: def _moved(values: Mapping[str, Any], device: torch.device) -> dict[str, Any]: - """The mapping with every tensor in it on ``device``, the rest untouched.""" - return { - name: value.to(device) if torch.is_tensor(value) else value - for name, value in values.items() - } + """The mapping with every array in it on ``device``, the rest untouched. + + Off the host, a list or tuple of numbers arrives as the tensor + :func:`as_torch` makes of it; on the host it stays as the caller wrote it. + """ + return {name: _on(value, device) for name, value in values.items()} + + +def _on(value: Any, device: torch.device) -> Any: + if torch.is_tensor(value): + return value.to(device) + numbers_given = ( + isinstance(value, (list, tuple)) + and len(value) > 0 + and all(isinstance(entry, numbers.Real) for entry in value) + ) + if numbers_given and device.type != "cpu": + return as_torch(value).to(device) + return value diff --git a/tests/recon/test_operator.py b/tests/recon/test_operator.py index bb36521d..0c42ffa5 100644 --- a/tests/recon/test_operator.py +++ b/tests/recon/test_operator.py @@ -283,6 +283,20 @@ def test_a_policy_changes_where_the_work_runs_and_nothing_else( torch.testing.assert_close(theirs, ours, atol=1e-5, rtol=1e-5) +@pytest.mark.parametrize("echo_times", [TE_MS.tolist(), tuple(TE_MS.tolist())]) +def test_maps_on_a_card_meet_echo_times_written_as_numbers_there(echo_times) -> None: + """A protocol given as plain numbers travels to the card with the maps.""" + if not torch.cuda.is_available(): + pytest.skip("CUDA is unavailable") + operator = ModelOperator(MultiEchoSimulator(TE=echo_times), "T2", bounds=BOUND) + x = operator.initial((7,), T2=80.0) + + there = operator.A(x.cuda()) + + assert there.device.type == "cuda" + torch.testing.assert_close(there.cpu(), operator.A(x), atol=1e-5, rtol=1e-5) + + def test_streaming_a_volume_too_big_for_the_budget(operator) -> None: """The chunking is the policy's, and the seams do not show.""" x = operator.initial((20_000,), T2=80.0)