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
34 changes: 31 additions & 3 deletions src/blochsim/_derivative.py
Original file line number Diff line number Diff line change
@@ -1,22 +1,50 @@
"""Directional derivatives taken in reverse mode.
"""Directional derivatives, and PyTorch's forward mode made ready quietly.

PyTorch's forward mode loads, the first time it makes a dual tensor in a
process, decompositions it registers through TorchScript. What the package
works out about its own structure -- how a packing moves with its arguments,
how a transition table moves with its flip -- is therefore taken in reverse
mode, so that a simulation and its gradients compile nothing at run time.
mode, so that a simulation and its gradients compile nothing at run time. The
derivatives that are forward mode by nature -- a Jacobian, a linearised
operator -- load those decompositions through :func:`forward_mode`, which
keeps PyTorch's deprecation of ``torch.jit.script`` out of the output.
"""

from __future__ import annotations

__all__ = ["directional_derivatives"]
__all__ = ["directional_derivatives", "forward_mode"]

import importlib
import os
import warnings
from collections.abc import Callable, Sequence
from functools import cache
from typing import Any

import torch


@cache
def forward_mode() -> None:
"""Load the decompositions PyTorch's forward mode registers, without its warning.

They are what PyTorch would load at its first dual tensor, under the switch
it reads for that: scripting them is PyTorch's own call, which warns that
``torch.jit.script`` is deprecated. Loaded once, here, the warning is
silenced for that load alone.
"""
if os.environ.get("PYTORCH_JIT", "1") != "1" or not __debug__:
return
with warnings.catch_warnings():
warnings.filterwarnings(
"ignore", message=r"`torch\.jit\.script`", category=DeprecationWarning
)
try:
importlib.import_module("torch._decomp.decompositions_for_jvp")
except ImportError:
return


def directional_derivatives(
function: Callable[..., Any],
primals: Sequence[torch.Tensor],
Expand Down
2 changes: 2 additions & 0 deletions src/blochsim/model/_signal.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@

import torch

from .._derivative import forward_mode
from ..sequence._array import as_torch, brought, is_array, like
from ..sequence._parameters import PUBLIC_PROPERTIES

Expand Down Expand Up @@ -228,6 +229,7 @@ def along(*inputs: torch.Tensor) -> torch.Tensor:
declared = {**rest, **dict(zip(names, inputs, strict=True))}
return self._shaped(self.evaluate(declared, **sequence), batch)

forward_mode()
columns = []
signal = None
for name in names:
Expand Down
2 changes: 2 additions & 0 deletions src/blochsim/recon/_operator.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@
import torch

from .._bounds import bound_of, to_free, to_natural, widen
from .._derivative import forward_mode
from .._execution import PER_VOXEL_CROSSOVER, per_voxel

#: What the amplitude occupies, when it is carried: real part then imaginary.
Expand Down Expand Up @@ -307,6 +308,7 @@ def A_jvp(self, x: torch.Tensor, d: torch.Tensor) -> torch.Tensor:
torch.Tensor
``(..., contrasts)``, complex.
"""
forward_mode()
return self._voxelwise(
"jvp",
(x, d),
Expand Down
Loading