"""RF coils of the virtual scanner, with the sensitivities of BART's coil models or of each coil's field maps."""
from __future__ import annotations
__all__ = ["COILS", "Coil", "coils"]
import functools
import json
import tempfile
from dataclasses import dataclass
from pathlib import Path
from typing import NamedTuple
import numpy as np
from .._accelerators import require
#: Side, in m, of the cube centred on the isocentre that the unit field of view
#: of BART's coil models is taken to span.
MODEL_FOV = 0.256
#: Samples per axis of the grid a model is evaluated on and interpolated from.
_SAMPLES = 32
#: Largest relative difference allowed between the frequency a coil's maps
#: were solved at and the Larmor frequency of the magnet they are used in.
_FREQUENCY_TOLERANCE = 0.01
class _Grid(NamedTuple):
"""Sensitivities ``(channels, z, y, x)``; sample ``i`` along an axis lies ``(i - centre) * step`` m from the isocentre."""
values: np.ndarray
centre: np.ndarray
step: np.ndarray
[docs]
@dataclass(frozen=True)
class Coil:
"""The RF coils of an exam, fixed in the physical frame: the coil pulses are transmitted on and the coil the signal is received by.
``name`` is ``transmit/receive``, or one coil's name where it does both.
Each side is BART's coil model (``HEAD_2D_8CH`` or ``HEAD_3D_64CH``) and
the number of its channels the coil takes, the first ones, as ``phantom
-S`` selects them, or the file of the coil's field maps :func:`coils`
reads; a side without either is one channel of unit sensitivity
everywhere, as a body coil is taken to be. A model is sampled by bartorch
on first use, which the ``coils`` extra installs, over :data:`MODEL_FOV`
along the physical axes; ``HEAD_2D_8CH`` is constant along z. Either is
interpolated linearly.
Receive sensitivities are scaled to a root sum of squares of 1 at the
isocentre, and transmit sensitivities so that the :attr:`default_shim`
plays a pulse at its nominal amplitude at the isocentre. A model's
transmit sensitivities are the complex conjugates of its receive
sensitivities, the quasi-static limit of reciprocity.
"""
name: str
transmit_model: tuple[str, int] | Path | None = None
receive_model: tuple[str, int] | Path | None = None
@property
def transmit_channels(self) -> int:
"""Channels a pulse is played on."""
return _channels(self.transmit_model)
@property
def receive_channels(self) -> int:
"""Coils a readout is received by."""
return _channels(self.receive_model)
@property
def default_shim(self) -> np.ndarray | None:
"""Unit channel weights that bring every transmit channel into phase at the isocentre; None for one channel.
A single-channel pulse played without an RF shim is played on every
channel with these weights. They are an RF shim's weights, which scale
the waveform in Pulseq's convention, the conjugate of the field.
"""
if self.transmit_model is None:
return None
return np.exp(1j * np.angle(_isocentre(_transmit(self.transmit_model))))
[docs]
def transmit(self, points: np.ndarray) -> np.ndarray | None:
"""Return each channel's transmit sensitivity at ``(n, 3)`` physical points, in m: ``(n, channels)``; None for one uniform channel."""
if self.transmit_model is None:
return None
return _interpolated(_transmit(self.transmit_model), points)
[docs]
def receive(self, points: np.ndarray) -> np.ndarray | None:
"""Return each coil's receive sensitivity at ``(n, 3)`` physical points, in m: ``(n, coils)``; None for one uniform coil."""
if self.receive_model is None:
return None
return _interpolated(_receive(self.receive_model), points)
[docs]
def limits(self) -> dict[str, str]:
"""Return the VOP entries of the check limits a design is made under in this coil; none unless its transmit maps have VOPs beside them.
``vop_file`` is ``<coil>_vops.npz`` beside the transmit maps
``<coil>.npz``; ``vop_drive_per_hz`` the drive of every channel, in
the unit of the maps and the VOPs, per Hz of a pulse's amplitude,
with which the channels' rotating fields, half their ``minus``, add
up to the pulse's nominal amplitude at the isocentre through the
:attr:`default_shim`; and ``vop_default_shim`` that shim, a magnitude
and a phase per channel.
"""
maps = self.transmit_model
if not isinstance(maps, Path):
return {}
vops = maps.with_name(f"{maps.stem}_vops.npz")
if not vops.is_file():
return {}
phases = np.angle(self.default_shim).tolist()
return {
"vop_file": str(vops),
"vop_drive_per_hz": repr(_drive_per_hz(maps)),
"vop_default_shim": " ".join(f"1.0 {phase!r}" for phase in phases),
}
COILS = {
coil.name: coil
for coil in (
Coil("body"),
Coil("body/head48", receive_model=("HEAD_3D_64CH", 48)),
Coil("head8/head32", ("HEAD_2D_8CH", 8), ("HEAD_3D_64CH", 32)),
)
}
def coils(
fields: Path | str | None = None, *, field_t: float | None = None
) -> dict[str, Coil]:
"""Return the virtual scanner's coils by name: :data:`COILS`, or the same coils with the field maps in ``fields``.
Each coil of a name has its maps in ``<coil>.npz``, as mariepy's
``maps.write`` writes them: the circular components ``plus`` and
``minus``, mu0 (Hx + j Hy) and mu0 (Hx - j Hy) with the time dependence
exp(+j omega t), of every channel's field over a body in the physical
frame, zero outside the body's mask. In a static field along +z, the
field a channel transmits is the complex conjugate of its ``minus``, and
what it receives is weighted by the complex conjugate of its ``plus``.
Voxels outside the mask take the value of the nearest voxel inside.
Raises
------
FileNotFoundError
If ``fields`` lacks a coil's maps.
ValueError
If maps were solved at a frequency other than the Larmor frequency of
``field_t``, in T.
"""
if fields is None:
return COILS
directory = Path(fields)
mapped = {}
for name in COILS:
transmit, _, receive = name.partition("/")
mapped[name] = Coil(
name,
directory / f"{transmit}.npz",
directory / f"{receive or transmit}.npz",
)
if field_t is not None:
import pypulseqpp as pp
larmor = pp.Opts().gamma * field_t
sides = {side for name in COILS for side in name.split("/")}
for path in sorted(directory / f"{side}.npz" for side in sides):
solved = float(_metadata(path)["frequency_hz"])
if abs(solved - larmor) > _FREQUENCY_TOLERANCE * larmor:
raise ValueError(
f"{path} was solved at {solved / 1e6:.2f} MHz, and the magnet's "
f"Larmor frequency is {larmor / 1e6:.2f} MHz"
)
return mapped
def _channels(side: tuple[str, int] | Path | None) -> int:
if side is None:
return 1
if isinstance(side, tuple):
return side[1]
return len(_metadata(side)["channels"])
@functools.cache
def _receive(side: tuple[str, int] | Path) -> _Grid:
grid = _model(*side) if isinstance(side, tuple) else _map(side, "plus")
return _scaled(grid, np.sqrt(np.sum(np.abs(_isocentre(grid)) ** 2)))
@functools.cache
def _transmit(side: tuple[str, int] | Path) -> _Grid:
if isinstance(side, tuple):
grid = _model(*side)
grid = grid._replace(values=np.conj(grid.values))
else:
grid = _map(side, "minus")
return _scaled(grid, np.sum(np.abs(_isocentre(grid))))
@functools.cache
def _drive_per_hz(maps: Path) -> float:
import pypulseqpp as pp
minus = np.sum(np.abs(_isocentre(_map(maps, "minus"))))
return 2.0 / (pp.Opts().gamma * float(minus))
def _scaled(grid: _Grid, norm: float) -> _Grid:
return grid._replace(values=grid.values / np.float32(norm))
def _model(model: str, channels: int) -> _Grid:
maps = _sampled(model, channels)
sizes = np.array(maps.shape[:0:-1])
return _Grid(maps, sizes // 2, MODEL_FOV / sizes)
def _sampled(model: str, channels: int) -> np.ndarray:
"""Return the first ``channels`` sensitivities of BART's ``model``: ``(channels, z, y, x)``, z of size 1 for a 2D model.
Sample ``i`` along an axis lies at ``(i - n // 2) / n`` of the model's
field of view, as BART's ``phantom`` places it.
"""
try:
from bartorch.tools import phantom
except ImportError as error:
raise ImportError(
"BART's coil models are sampled by bartorch: pip install 'pulserver[coils]'"
) from error
three = model == "HEAD_3D_64CH"
maps = phantom((_SAMPLES,) * (3 if three else 2), S=channels, coil=model)
return np.asarray(maps.cpu().numpy(), dtype=np.complex64).reshape(
channels, -1, _SAMPLES, _SAMPLES
)
def _map(path: Path, component: str) -> _Grid:
"""Return the complex conjugate of a map file's ``plus`` or ``minus``, each voxel outside its mask given the value of the nearest voxel inside."""
from scipy import ndimage
with np.load(path, allow_pickle=False) as archive:
values, mask = archive[component], archive["mask"]
nearest = ndimage.distance_transform_edt(
~mask, return_distances=False, return_indices=True
)
values = np.conj(values[(slice(None), *nearest)])
metadata = _metadata(path)
step = np.full(3, float(metadata["resolution"]))
return _Grid(
np.ascontiguousarray(values.transpose(0, 3, 2, 1)),
-np.asarray(metadata["origin"], dtype=float) / step,
step,
)
@functools.cache
def _metadata(path: Path) -> dict:
with np.load(path, allow_pickle=False) as archive:
return json.loads(str(archive["metadata"]))
def _isocentre(grid: _Grid) -> np.ndarray:
return _trilinear(grid, np.zeros((1, 3)))[0]
def _interpolated(grid: _Grid, points: np.ndarray) -> np.ndarray:
"""Return ``grid`` interpolated trilinearly at ``(n, 3)`` physical points, as :func:`_trilinear` interpolates it; the edge value beyond it.
The result is complex128, as :class:`Isochromats` takes it,
in a temporary file mapped into memory: its pages belong to the file,
which the operating system writes back rather than holding in the
process's memory. ``TMPDIR`` names where the file is made.
"""
points = np.ascontiguousarray(np.asarray(points, dtype=float).reshape(-1, 3))
out = _mapped((len(points), grid.values.shape[0]))
require("bloch").trilinear(
grid.values,
np.asarray(grid.centre, dtype=float),
np.asarray(grid.step, dtype=float),
points,
out,
)
return out
def _trilinear(grid: _Grid, points: np.ndarray) -> np.ndarray:
table = grid.values.reshape(grid.values.shape[0], -1).T
corners, weights = [], []
for axis, size in zip((2, 1, 0), grid.values.shape[1:], strict=True):
index = np.clip(
points[:, axis] / grid.step[axis] + grid.centre[axis], 0, size - 1
)
low = np.minimum(np.floor(index).astype(np.intp), max(size - 2, 0))
corners.append((low, np.minimum(low + 1, size - 1)))
weights.append((1.0 - (index - low), index - low))
total = np.zeros((len(points), grid.values.shape[0]), dtype=np.complex64)
for z in (0, 1):
for y in (0, 1):
for x in (0, 1):
flat = np.ravel_multi_index(
(corners[0][z], corners[1][y], corners[2][x]), grid.values.shape[1:]
)
weight = weights[0][z] * weights[1][y] * weights[2][x]
total += weight.astype(np.float32)[:, None] * table[flat]
return total
def _mapped(shape: tuple[int, int]) -> np.ndarray:
"""Return a complex128 array of ``shape`` in an unlinked temporary file, which is freed with the array."""
if 0 in shape:
return np.empty(shape, dtype=np.complex128)
with tempfile.TemporaryFile() as file:
return np.memmap(file, dtype=np.complex128, mode="w+", shape=shape)