Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 20 additions & 5 deletions src/torchsim/model/_signal.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
14 changes: 14 additions & 0 deletions tests/recon/test_operator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading