diff --git a/src/blochsim/_derivative.py b/src/blochsim/_derivative.py index 9553b38b..51ea16c2 100644 --- a/src/blochsim/_derivative.py +++ b/src/blochsim/_derivative.py @@ -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], diff --git a/src/blochsim/model/_signal.py b/src/blochsim/model/_signal.py index 98967169..5ce8144c 100644 --- a/src/blochsim/model/_signal.py +++ b/src/blochsim/model/_signal.py @@ -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 @@ -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: diff --git a/src/blochsim/recon/_operator.py b/src/blochsim/recon/_operator.py index acced5ca..aa213da2 100644 --- a/src/blochsim/recon/_operator.py +++ b/src/blochsim/recon/_operator.py @@ -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. @@ -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),