diff --git a/src/frontx/_boltzmann.py b/src/frontx/_boltzmann.py index 706c22f..3369207 100644 --- a/src/frontx/_boltzmann.py +++ b/src/frontx/_boltzmann.py @@ -1,7 +1,7 @@ from abc import abstractmethod from collections.abc import Callable from functools import wraps -from typing import Any, Protocol, TypeVar, overload +from typing import Protocol, TypeVar, overload import diffrax import equinox as eqx @@ -14,13 +14,19 @@ def ode( D: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], + [ + float + | jax.Array + | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]] + ], + float + | jax.Array + | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]], ], ) -> diffrax.ODETerm[jax.Array]: - @diffrax.ODETerm[jax.Array] + @diffrax.ODETerm[jax.Array] # ty: ignore[invalid-argument-type] def term( - o: float | jax.Array | np.ndarray[Any, Any], + o: jax.Array, y: jax.Array, args: None, ) -> jax.Array: @@ -42,41 +48,68 @@ class _BoltzmannTransformed(Protocol): @overload def __call__( self, - r: float | jax.Array | np.ndarray[Any, Any], - t: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: ... + r: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + t: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> jax.Array: ... @overload def __call__( - self, o: float | jax.Array | np.ndarray[Any, Any] - ) -> float | jax.Array | np.ndarray[Any, Any]: ... + self, + o: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> jax.Array: ... def boltzmannmethod( meth: Callable[ - [_Self, float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], + [ + _Self, + float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ], + float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], ], /, ) -> _BoltzmannTransformed: @overload def boltzmann_wrapper( self: _Self, - r: float | jax.Array | np.ndarray[Any, Any], - t: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: ... + r: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + t: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> ( + float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + ): ... @overload def boltzmann_wrapper( - self: _Self, o: float | jax.Array | np.ndarray[Any, Any] - ) -> float | jax.Array | np.ndarray[Any, Any]: ... + self: _Self, + o: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> ( + float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + ): ... @wraps(meth) def boltzmann_wrapper( self: _Self, - *args: float | jax.Array | np.ndarray[Any, Any], - **kwargs: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: + *args: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + **kwargs: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: match args, kwargs: case (o,), {} if not kwargs: pass @@ -98,61 +131,86 @@ def boltzmann_wrapper( class AbstractSolution(eqx.Module): D: eqx.AbstractVar[ Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], + [ + float + | jax.Array + | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]] + ], + float + | jax.Array + | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]], ] ] oi: eqx.AbstractVar[float] @property - def b(self) -> float | jax.Array | np.ndarray[Any, Any]: + def b(self) -> jax.Array: return self(0.0) @property - def d_dob(self) -> float | jax.Array | np.ndarray[Any, Any]: + def d_dob(self) -> jax.Array: return self.d_do(0.0) @property - def i(self) -> float | jax.Array | np.ndarray[Any, Any]: + def i(self) -> jax.Array: return self(self.oi) @abstractmethod @boltzmannmethod def __call__( self, - o: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: + o: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: raise NotImplementedError @boltzmannmethod def d_do( self, - o: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: + o: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: return vmap(jax.grad(self))(o) def d_dr( self, - r: float | jax.Array | np.ndarray[Any, Any], - t: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: + r: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + t: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> jax.Array: return self.d_do(r, t) / jnp.sqrt(t) def d_dt( self, - r: float | jax.Array | np.ndarray[Any, Any], - t: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: + r: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + t: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> jax.Array: return -r / (jnp.sqrt(t) * 2 * t) * self.d_do(r, t) def flux( self, - r: float | jax.Array | np.ndarray[Any, Any], - t: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: - return -self.D(self(r, t)) * self.d_dr(r, t) + r: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + t: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> jax.Array: + return -self.D(self(r, t)) * self.d_dr(r, t) # ty: ignore[invalid-return-type] def sorptivity( - self, o: float | jax.Array | np.ndarray[Any, Any] = 0.0 - ) -> float | jax.Array | np.ndarray[Any, Any]: - return -2 * self.D(self(o)) * self.d_do(o) + self, + o: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] = 0.0, + ) -> jax.Array: + return -2 * self.D(self(o)) * self.d_do(o) # ty: ignore[invalid-return-type] diff --git a/src/frontx/_forward.py b/src/frontx/_forward.py index 71de321..675f9cb 100644 --- a/src/frontx/_forward.py +++ b/src/frontx/_forward.py @@ -1,5 +1,5 @@ from collections.abc import Callable -from typing import Any +from typing import cast import diffrax import equinox as eqx @@ -13,64 +13,82 @@ RESULTS = diffrax.RESULTS +_Diffusivity = Callable[ + [ + float + | jax.Array + | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]] + ], + float | jax.Array | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]], +] +_DiffusivityInput = ( + _Diffusivity + | Callable[ + [ + float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + ], + float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ] +) + class Solution(AbstractSolution): _sol: diffrax.Solution result: RESULTS - D: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], - ] + D: _Diffusivity @boltzmannmethod def __call__( self, - o: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: - return vmap(self._sol.evaluate)(jnp.clip(o, 0, self.oi))[..., 0] + o: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: + return vmap(self._sol.evaluate)(jnp.clip(o, 0, self.oi))[..., 0] # ty: ignore[not-subscriptable] @boltzmannmethod def d_do( self, - o: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: - return vmap(self._sol.evaluate)(jnp.clip(o, 0, self.oi))[..., 1] + o: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: + return vmap(self._sol.evaluate)(jnp.clip(o, 0, self.oi))[..., 1] # ty: ignore[not-subscriptable] @property - def oi(self) -> float: + def oi(self) -> jax.Array: assert self._sol.ts is not None return self._sol.ts[-1] @property - def i(self) -> float: + def i(self) -> jax.Array: assert self._sol.ys is not None return self._sol.ys[-1, 0] @property - def b(self) -> float: + def b(self) -> jax.Array: assert self._sol.ys is not None return self._sol.ys[0, 0] @property - def d_dob(self) -> float: + def d_dob(self) -> jax.Array: assert self._sol.ys is not None return self._sol.ys[0, 1] @eqx.filter_jit def solve( - D: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], - ], + D: _DiffusivityInput, *, - b: float, - i: float, + b: float | jax.Array, + i: float | jax.Array, itol: float = 1e-3, max_steps: int = 100, throw: bool = True, ) -> Solution: - term = ode(D) + term = ode(cast(_Diffusivity, D)) direction = jnp.sign(i - b) @diffrax.Event @@ -103,7 +121,7 @@ def shoot( root: optx.Solution = optx.root_find( shoot, - solver=optx.Bisection(rtol=jnp.inf, atol=itol, expand_if_necessary=True), # ty: ignore[missing-argument] + solver=optx.Bisection(rtol=jnp.inf, atol=itol, expand_if_necessary=True), y0=0, max_steps=max_steps, has_aux=True, @@ -112,11 +130,11 @@ def shoot( ) return Solution( - root.aux, - RESULTS.where( + _sol=root.aux, + result=RESULTS.where( root.result == optx.RESULTS.successful, RESULTS.successful, RESULTS.max_steps_reached, ), - D, # ty: ignore[invalid-argument-type] + D=cast(_Diffusivity, D), ) # ty: ignore[missing-argument] diff --git a/src/frontx/_inverse/fit.py b/src/frontx/_inverse/fit.py index e21c50f..f21383f 100644 --- a/src/frontx/_inverse/fit.py +++ b/src/frontx/_inverse/fit.py @@ -1,5 +1,5 @@ from collections.abc import Callable -from typing import Any, Literal +from typing import Literal import equinox as eqx import jax @@ -46,10 +46,12 @@ def with_sorptivity( @staticmethod def fitting_data( original: AbstractSolution, - o: jax.Array | np.ndarray[Any, Any], - theta: jax.Array | np.ndarray[Any, Any], + o: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + theta: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - sigma: float | jax.Array | np.ndarray[Any, Any] = 1, + sigma: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] = 1, *, throw: bool = True, ) -> "ScaledSolution": @@ -84,41 +86,57 @@ def residuals( @boltzmannmethod def __call__( self, - o: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: + o: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> jax.Array: return self.original(o / jnp.sqrt(self.D0)) @boltzmannmethod def d_do( self, - o: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: + o: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> jax.Array: return self.original.d_do(o / jnp.sqrt(self.D0)) / jnp.sqrt(self.D0) def D( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> ( + float + | jax.Array + | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]] + ): return self.original.D(theta) * self.D0 @property - def oi(self) -> float | jax.Array: + def oi(self) -> jax.Array: return self.original.oi * jnp.sqrt(self.D0) def fit( D: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], + [ + float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + ], + float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], ], - o: jax.Array | np.ndarray[Any, Any], - theta: jax.Array | np.ndarray[Any, Any], + o: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + theta: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - sigma: float | jax.Array | np.ndarray[Any, Any] = 1, + sigma: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] = 1, *, - i: float, - b: float, + i: float | jax.Array, + b: float | jax.Array, fit_D0: Literal["data", "sorptivity"] | None = "data", max_steps: int = 15, ) -> ScaledSolution | Solution: @@ -127,8 +145,14 @@ def fit( def candidate( D: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], + [ + float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + ], + float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], ], ) -> ScaledSolution | Solution: sol = solve(D, i=i, b=b, throw=False) diff --git a/src/frontx/_inverse/interpolated.py b/src/frontx/_inverse/interpolated.py index 45b3eb1..a2defbb 100644 --- a/src/frontx/_inverse/interpolated.py +++ b/src/frontx/_inverse/interpolated.py @@ -1,5 +1,3 @@ -from typing import Any - import jax import jax.numpy as jnp import numpy as np @@ -17,8 +15,8 @@ class InterpolatedSolution(AbstractSolution): def __init__( self, - o: jax.Array | np.ndarray[Any, Any], - theta: jax.Array | np.ndarray[Any, Any], + o: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + theta: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, *, b: float | jax.Array | None = None, @@ -52,22 +50,29 @@ def __init__( @boltzmannmethod def __call__( self, - o: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: + o: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> jax.Array: return self._sol(o) def D( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> jax.Array: Iodtheta = self._Iodtheta(theta) - self._c do_dtheta = self._do_dtheta(theta) return jnp.squeeze(-(do_dtheta * Iodtheta) / 2) def sorptivity( - self, o: float | jax.Array | np.ndarray[Any, Any] = 0 - ) -> float | jax.Array | np.ndarray[Any, Any]: + self, + o: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] = 0.0, + ) -> jax.Array: Ithetado = self._sol.antiderivative() return (Ithetado(self.oi) - Ithetado(o)) - self.i * (self.oi - o) diff --git a/src/frontx/_inverse/param.py b/src/frontx/_inverse/param.py index b18cdd7..54f2479 100644 --- a/src/frontx/_inverse/param.py +++ b/src/frontx/_inverse/param.py @@ -1,5 +1,5 @@ from collections.abc import Callable, Sequence -from typing import Any, TypeVar, overload +from typing import TypeVar, overload import jax import jax.numpy as jnp @@ -9,14 +9,14 @@ class Param(pmx.Param[float]): - min: float | None - max: float | None + min: float | jax.Array | None + max: float | jax.Array | None def __init__( self, value: float | jax.Array | None = None, - min: float | None = None, - max: float | None = None, + min: float | jax.Array | None = None, + max: float | jax.Array | None = None, ) -> None: if min is not None and max is not None: if value is None: @@ -68,7 +68,11 @@ def collect(obj: object) -> None: def set_param_values( - pytree: _T, values: jax.Array | np.ndarray[Any, Any] | Sequence[float], / + pytree: _T, + values: jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + | Sequence[float], + /, ) -> _T: i = 0 @@ -92,7 +96,7 @@ def replace(obj: Param | _O) -> Param | _O: def de_fit( candidate: Callable[[_T], _O], - cost: Callable[[_O], float], + cost: Callable[[_O], float | jax.Array], /, initial: _T, *, diff --git a/src/frontx/_inverse/sorptivity.py b/src/frontx/_inverse/sorptivity.py index c282cea..ce1e80a 100644 --- a/src/frontx/_inverse/sorptivity.py +++ b/src/frontx/_inverse/sorptivity.py @@ -1,17 +1,15 @@ -from typing import Any - import jax import jax.numpy as jnp import numpy as np def sorptivity( - o: jax.Array | np.ndarray[Any, Any], - theta: jax.Array | np.ndarray[Any, Any], + o: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + theta: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, *, - b: float, - i: float, + b: float | jax.Array, + i: float | jax.Array, ) -> jax.Array: o = jnp.insert(o, 0, 0) theta = jnp.insert(theta, 0, b) diff --git a/src/frontx/_util.py b/src/frontx/_util.py index 97af792..1475c11 100644 --- a/src/frontx/_util.py +++ b/src/frontx/_util.py @@ -1,52 +1,34 @@ from collections.abc import Callable from functools import wraps -from typing import Any, overload import jax import jax.numpy as jnp import numpy as np -@overload def vmap( func: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - jax.Array | np.ndarray[Any, Any], + [ + float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + ], + float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], ], /, ) -> Callable[ - [float | jax.Array | np.ndarray[Any, Any]], jax.Array | np.ndarray[Any, Any] -]: ... - - -@overload -def vmap( - func: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], - ], - /, -) -> Callable[ - [float | jax.Array | np.ndarray[Any, Any]], float | jax.Array | np.ndarray[Any, Any] -]: ... - - -def vmap( - func: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], - ], - /, -) -> Callable[ - [float | jax.Array | np.ndarray[Any, Any]], float | jax.Array | np.ndarray[Any, Any] + [float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]], + float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], ]: vfunc = jax.vmap(func) @wraps(func) def vmap_wrapper( - x: float | jax.Array | np.ndarray[Any, Any], + x: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: if jnp.ndim(x) == 0: return func(x) diff --git a/src/frontx/examples/neural/exact.py b/src/frontx/examples/neural/exact.py index c429c06..a0b154d 100755 --- a/src/frontx/examples/neural/exact.py +++ b/src/frontx/examples/neural/exact.py @@ -14,8 +14,6 @@ python -m frontx.examples.exacti """ -from typing import Any - import jax import jax.numpy as jnp import matplotlib.pyplot as plt @@ -27,7 +25,11 @@ jax.config.update("jax_enable_x64", True) -def D(theta: float | jax.Array | np.ndarray[Any, Any]) -> float | jax.Array: +def D( + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], +) -> float | jax.Array: """Custom diffusivity used by the PINN. Args: diff --git a/src/frontx/examples/neural/grenoblesand.py b/src/frontx/examples/neural/grenoblesand.py index 75e0cd7..4c007ad 100755 --- a/src/frontx/examples/neural/grenoblesand.py +++ b/src/frontx/examples/neural/grenoblesand.py @@ -54,7 +54,7 @@ def run() -> None: b=theta_s - 1e-7, ) - print(f"Ks={sol.D.Ks.value}, m={sol.D.m.value}") + print(f"Ks={sol.D.Ks.value}, m={sol.D.m.value}") # ty: ignore[unresolved-attribute] # Display o_display = np.linspace(0, ref.oi * 1.5, 500) diff --git a/src/frontx/finite.py b/src/frontx/finite.py index b88e179..570cfc7 100644 --- a/src/frontx/finite.py +++ b/src/frontx/finite.py @@ -12,7 +12,6 @@ """ from collections.abc import Callable -from typing import Any import diffrax import equinox as eqx @@ -39,17 +38,25 @@ class Solution(eqx.Module): """ D: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], + [ + float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + ], + float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], ] - r1: float + r1: float | jax.Array _sol: diffrax.Solution def __call__( self, - r: float | jax.Array | np.ndarray[Any, Any], - t: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: + r: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + t: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> jax.Array: """Evaluate the simulated field at coordinates ``(r, t)``. Dense time output is queried from the solver and linearly interpolated @@ -68,9 +75,9 @@ def __call__( return jnp.interp(r, jnp.linspace(0, self.r1, theta.size), theta) @property - def t1(self) -> float | jax.Array | np.ndarray[Any, Any]: + def t1(self) -> float | jax.Array: """Final integration time.""" - return self._sol.t1 + return self._sol.t1 # ty: ignore[invalid-return-type] @property def result(self) -> RESULTS: @@ -81,15 +88,19 @@ def result(self) -> RESULTS: @eqx.filter_jit def solve( D: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], + [ + float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + ], + float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], ], - r1: float, - t1: float, + r1: float | jax.Array, + t1: float | jax.Array, *, - i: jax.Array | np.ndarray[Any, Any], - b: float | None = None, - throw: bool = True, + i: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + b: float | jax.Array | None = None, + throw: bool | jax.Array = True, ) -> Solution: """Integrate the finite-difference diffusion model. @@ -126,8 +137,8 @@ def solve( @diffrax.ODETerm[jax.Array] def term( - _: float | jax.Array | np.ndarray[Any, Any], - theta: jax.Array, + _: object, + theta: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], args: None, ) -> jax.Array: D_ = jnp.asarray(D(theta)) @@ -166,19 +177,25 @@ def term( def fit( D: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], + [ + float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + ], + float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], ], - r1: float, - t1: float, - r: jax.Array | np.ndarray[Any, Any], - t: jax.Array | np.ndarray[Any, Any], - theta: jax.Array | np.ndarray[Any, Any], + r1: float | jax.Array, + t1: float | jax.Array, + r: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + t: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + theta: jax.Array | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]], /, - sigma: float | jax.Array | np.ndarray[Any, Any] = 1, + sigma: float + | jax.Array + | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]] = 1, *, - i: jax.Array | np.ndarray[Any, Any], - b: float | None = None, + i: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + b: float | jax.Array | None = None, max_steps: int = 15, ) -> Solution: """Fit the finite-difference model to spatio-temporal observations. @@ -213,13 +230,19 @@ def fit( def candidate( D: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], + [ + float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + ], + float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], ], ) -> Solution: return solve(D, r1, t1, i=i, b=b) - def cost(sol: Solution) -> float: + def cost(sol: Solution) -> jax.Array: return jax.lax.cond( sol.result == RESULTS.successful, lambda: jnp.mean(((sol(r, t[:, jnp.newaxis]) - theta) / sigma) ** 2), diff --git a/src/frontx/models.py b/src/frontx/models.py index 4ce5817..1faf9a8 100644 --- a/src/frontx/models.py +++ b/src/frontx/models.py @@ -1,7 +1,6 @@ """Moisture diffusivity models.""" from abc import abstractmethod -from typing import Any import equinox as eqx import jax @@ -16,9 +15,11 @@ class _MoistureDiffusivityModel(eqx.Module): def _Se( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: return (theta - self.theta_range[0]) / ( self.theta_range[1] - self.theta_range[0] ) @@ -26,9 +27,11 @@ def _Se( @abstractmethod def __call__( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: raise NotImplementedError @@ -53,28 +56,32 @@ class LETd(_MoistureDiffusivityModel): theta_range: Tuple ``(theta_r, theta_s)`` used to compute ``Se``. """ - L: float | Param - E: float | Param - T: float | Param + L: float | Param # ty: ignore[dataclass-field-order] + E: float | Param # ty: ignore[dataclass-field-order] + T: float | Param # ty: ignore[dataclass-field-order] Dwt: float | Param = 1 theta_range: tuple[float | Param, float | Param] = (0, 1) def __call__( - self, theta: float | jax.Array | np.ndarray[Any, Any] - ) -> float | jax.Array | np.ndarray[Any, Any]: + self, + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + /, + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: Se = (theta - self.theta_range[0]) / (self.theta_range[1] - self.theta_range[0]) return self.Dwt * Se**self.L / (Se**self.L + self.E * (1 - Se) ** self.T) class _RichardsModel(_MoistureDiffusivityModel): - Ks: eqx.AbstractVar[float | Param | None] - k: eqx.AbstractVar[float | Param | None] - g: eqx.AbstractVar[float | Param] - rho: eqx.AbstractVar[float | Param] - mu: eqx.AbstractVar[float | Param] + Ks: eqx.AbstractVar[float | jax.Array | Param | None] + k: eqx.AbstractVar[float | jax.Array | Param | None] + g: eqx.AbstractVar[float | jax.Array | Param] + rho: eqx.AbstractVar[float | jax.Array | Param] + mu: eqx.AbstractVar[float | jax.Array | Param] @property - def _Ks(self) -> float | Param | jax.Array: + def _Ks(self) -> float | jax.Array | Param: if self.Ks is None: if self.k is None: return 1 @@ -87,39 +94,49 @@ def _Ks(self) -> float | Param | jax.Array: def __call__( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: - return self._K(theta) / self._C(theta) + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: + return self._K(theta) / self._C(theta) # ty: ignore[invalid-return-type] def _C( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: return 1 / vmap(jax.grad(self._h))(theta) @abstractmethod def _h( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: raise NotImplementedError @abstractmethod def _kr( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: raise NotImplementedError def _K( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: return self._Ks * self._kr(theta) @@ -147,29 +164,33 @@ class BrooksAndCorey(_RichardsModel): theta_range: Tuple ``(theta_r, theta_s)`` for effective saturation. """ - n: float | Param - l: float | Param = 1 - Ks: float | Param | None = None - k: float | None = None - g: float | Param = 9.81 - rho: float | Param = 1e3 - mu: float | Param = 1e-3 - alpha: float | Param = 1 - theta_range: tuple[float | Param, float | Param] = (0, 1) + n: float | jax.Array | Param # ty: ignore[dataclass-field-order] + l: float | jax.Array | Param = 1 + Ks: float | jax.Array | Param | None = None + k: float | jax.Array | None = None + g: float | jax.Array | Param = 9.81 + rho: float | jax.Array | Param = 1e3 + mu: float | jax.Array | Param = 1e-3 + alpha: float | jax.Array | Param = 1 + theta_range: tuple[float | jax.Array | Param, float | jax.Array | Param] = (0, 1) def _h( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: Se = self._Se(theta) return -1 / (self.alpha * Se ** (1 / self.n)) def _kr( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: Se = self._Se(theta) return Se ** (2 / self.n + self.l + 2) @@ -203,16 +224,16 @@ class VanGenuchten(_RichardsModel): ValueError: If neither ``n`` nor ``m`` is provided. """ - n: float | Param | None = None - m: float | Param | None = None - l: float | Param = 0.5 - Ks: float | Param | None = None - k: float | Param | None = None - g: float | Param = 9.81 - rho: float | Param = 1e3 - mu: float | Param = 1e-3 - alpha: float | Param = 1 - theta_range: tuple[float | Param, float | Param] = (0, 1) + n: float | jax.Array | Param | None = None + m: float | jax.Array | Param | None = None + l: float | jax.Array | Param = 0.5 + Ks: float | jax.Array | Param | None = None + k: float | jax.Array | Param | None = None + g: float | jax.Array | Param = 9.81 + rho: float | jax.Array | Param = 1e3 + mu: float | jax.Array | Param = 1e-3 + alpha: float | jax.Array | Param = 1 + theta_range: tuple[float | jax.Array | Param, float | jax.Array | Param] = (0, 1) @property def _n(self) -> float | jax.Array | Param: @@ -238,17 +259,21 @@ def _m(self) -> float | jax.Array | Param: def _h( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: Se = self._Se(theta) return -((1 / (Se ** (1 / self._m)) - 1) ** (1 / self._n)) / self.alpha def _kr( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: Se = self._Se(theta) return Se**self.l * (1 - (1 - Se ** (1 / self._m)) ** self._m) ** 2 @@ -281,33 +306,37 @@ class LETxs(_RichardsModel): theta_range: Tuple ``(theta_r, theta_s)`` for effective saturation. """ - Lw: float | Param - Ew: float | Param - Tw: float | Param - Ls: float | Param - Es: float | Param - Ts: float | Param - Ks: float | Param | None = None - k: float | None = None - g: float | Param = 9.81 - rho: float | Param = 1e3 - mu: float | Param = 1e-3 - alpha: float | Param = 1 - theta_range: tuple[float | Param, float | Param] = (0, 1) + Lw: float | jax.Array | Param # ty: ignore[dataclass-field-order] + Ew: float | jax.Array | Param # ty: ignore[dataclass-field-order] + Tw: float | jax.Array | Param # ty: ignore[dataclass-field-order] + Ls: float | jax.Array | Param # ty: ignore[dataclass-field-order] + Es: float | jax.Array | Param # ty: ignore[dataclass-field-order] + Ts: float | jax.Array | Param # ty: ignore[dataclass-field-order] + Ks: float | jax.Array | Param | None = None + k: float | jax.Array | Param | None = None + g: float | jax.Array | Param = 9.81 + rho: float | jax.Array | Param = 1e3 + mu: float | jax.Array | Param = 1e-3 + alpha: float | jax.Array | Param = 1 + theta_range: tuple[float | jax.Array | Param, float | jax.Array | Param] = (0, 1) def _kr( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: Se = self._Se(theta) - return Se**self.Lw / (Se**self.Lw + self.Ew * (1 - Se) ** self.Tw) + return Se**self.Lw / (Se**self.Lw + self.Ew * (1 - Se) ** self.Tw) # ty: ignore[invalid-return-type] def _h( self, - theta: float | jax.Array | np.ndarray[Any, Any], + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - ) -> float | jax.Array | np.ndarray[Any, Any]: + ) -> float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]]: Se = self._Se(theta) return ( -((1 - Se) ** self.Ls / ((1 - Se) ** self.Ls + self.Es * Se**self.Ts)) diff --git a/src/frontx/neural.py b/src/frontx/neural.py index 9596ced..5108cb1 100644 --- a/src/frontx/neural.py +++ b/src/frontx/neural.py @@ -10,7 +10,7 @@ """ from collections.abc import Callable -from typing import Any +from typing import cast import equinox as eqx import jax @@ -24,8 +24,12 @@ class _PINN(eqx.Module): D: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], + [ + float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + ], + float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], ] _net: eqx.nn.MLP = eqx.field( default_factory=lambda: eqx.nn.MLP( @@ -39,20 +43,27 @@ class _PINN(eqx.Module): ) def __call__( - self, x: float | jax.Array | np.ndarray[Any, Any] - ) -> float | jax.Array | np.ndarray[Any, Any]: - return 2 * jax.nn.sigmoid(-x * jax.nn.softplus(vmap(self._net)(x))) # ty: ignore [no-matching-overload] + self, + x: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> jax.Array: + return 2 * jax.nn.sigmoid(-x * jax.nn.softplus(vmap(self._net)(x))) # ty: ignore[invalid-argument-type] def data_loss( self, - x_data: jax.Array | np.ndarray[Any, Any], - y_data: jax.Array | np.ndarray[Any, Any], + x_data: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + y_data: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - y_sigma: float | jax.Array | np.ndarray[Any, Any] = 1, + y_sigma: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] = 1, ) -> jax.Array: return jnp.mean(((self(x_data) - y_data) / (y_sigma)) ** 2) - def physics_residuals(self, /, *, i: float, b: float, oi: float) -> jax.Array: + def physics_residuals( + self, /, *, i: float | jax.Array, b: float | jax.Array, oi: float | jax.Array + ) -> jax.Array: x = jnp.linspace(0, 1, 501)[1:] lhs = -x / 2 * vmap(jax.grad(self))(x) @@ -68,7 +79,13 @@ def physics_residuals(self, /, *, i: float, b: float, oi: float) -> jax.Array: return lhs - rhs def physics_loss( - self, /, *, i: float, b: float, oi: float, residual_cutoff: float = jnp.inf + self, + /, + *, + i: float | jax.Array, + b: float | jax.Array, + oi: float | jax.Array, + residual_cutoff: float | jax.Array = jnp.inf, ) -> jax.Array: residuals = self.physics_residuals(i=i, b=b, oi=oi) @@ -97,18 +114,18 @@ class Solution(AbstractSolution): oi: Characteristic scale used to normalize inputs ``o`` (``x = o/oi``). """ - oi: float + oi: float | jax.Array _net: _PINN - _i: float - _b: float + _i: float | jax.Array + _b: float | jax.Array def __init__( self, net: _PINN, *, - i: float, - b: float, - oi: float, + i: float | jax.Array, + b: float | jax.Array, + oi: float | jax.Array, ) -> None: """Initialize a :class:`Solution`. @@ -128,8 +145,14 @@ def __init__( def D( self, ) -> Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], + [ + float + | jax.Array + | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]] + ], + float + | jax.Array + | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]], ]: """Return the diffusivity-like callable used in the physics term. @@ -138,13 +161,27 @@ def D( with the same broadcastable shape. It is the same function that was passed to :func:`fit` as ``D``. """ - return self._net.D + return cast( + Callable[ + [ + float + | jax.Array + | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]] + ], + float + | jax.Array + | np.ndarray[tuple[int, ...], np.dtype[np.floating | np.integer]], + ], + self._net.D, + ) @boltzmannmethod def __call__( self, - o: float | jax.Array | np.ndarray[Any, Any], - ) -> float | jax.Array | np.ndarray[Any, Any]: + o: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> jax.Array: """Evaluate the trained solution at input ``o``. The input is clipped to ``[0, oi]`` (after normalization) and the @@ -164,18 +201,25 @@ def __call__( @eqx.filter_jit def fit( D: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], + [ + float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + ], + float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], ], - o: jax.Array | np.ndarray[Any, Any], - theta: jax.Array | np.ndarray[Any, Any], + o: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + theta: jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], /, - sigma: float | jax.Array | np.ndarray[Any, Any] | None = None, + sigma: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + | None = None, *, - i: float, - b: float, - oi: float | None = None, - max_steps: int = 300_000, + i: float | jax.Array, + b: float | jax.Array, + oi: float | jax.Array | None = None, + max_steps: int | jax.Array = 300_000, ) -> Solution: """Train a PINN against data and physics, returning a callable solution. @@ -205,7 +249,7 @@ def fit( AssertionError: If internal normalization parameters are missing. """ if oi is None: - oi = o[-1] * 1.05 # ty: ignore [invalid-assignment] + oi = o[-1] * 1.05 assert oi is not None @@ -215,7 +259,7 @@ def fit( y_data = (theta - i) / (b - i) y_sigma = sigma / (b - i) if sigma is not None else 1 - initial_data_loss = net.data_loss(x_data, y_data, y_sigma=y_sigma) + initial_data_loss = net.data_loss(x_data, y_data, y_sigma=y_sigma) # ty: ignore[invalid-argument-type] trainable_net, static_net = eqx.partition(net, eqx.is_array) @@ -225,7 +269,10 @@ def fit( opt_state = optim.init(trainable_net) def loss( - trainable_net: _PINN, *, step: int, residual_cutoff: float = jnp.inf + trainable_net: _PINN, + *, + step: int | jax.Array, + residual_cutoff: float | jax.Array = jnp.inf, ) -> jax.Array: net = eqx.combine(trainable_net, static_net) @@ -234,7 +281,7 @@ def loss( i=i, b=b, oi=oi, residual_cutoff=residual_cutoff ) - data_loss = net.data_loss(x_data, y_data, y_sigma=y_sigma) + data_loss = net.data_loss(x_data, y_data, y_sigma=y_sigma) # ty: ignore[invalid-argument-type] lambda_ = initial_data_loss * 10 ** (-2 + step / 100_000) @@ -243,10 +290,10 @@ def loss( def train_step( trainable_net: _PINN, opt_state: optax.OptState, - step: int, - physics_loss: float, - residual_cutoff: float, - ) -> tuple[_PINN, optax.OptState, int, float, float]: + step: jax.Array, + physics_loss: jax.Array, + residual_cutoff: jax.Array, + ) -> tuple[_PINN, optax.OptState, jax.Array, jax.Array, jax.Array]: net = eqx.combine(trainable_net, static_net) assert oi is not None @@ -254,13 +301,13 @@ def train_step( spike_score = ( jnp.max(jnp.abs(residuals)) - jnp.percentile(jnp.abs(residuals), 99) ) / (jnp.median(jnp.abs(jnp.abs(residuals) - jnp.median(jnp.abs(residuals))))) - residual_cutoff = jax.lax.select( # ty: ignore [invalid-assignment] + residual_cutoff = jax.lax.select( (step >= 50_000) & (residual_cutoff == jnp.inf) & (spike_score > 200), jnp.mean(jnp.abs(residuals)), residual_cutoff, ) - physics_loss = net.physics_loss( # ty: ignore [invalid-assignment] + physics_loss = net.physics_loss( i=i, b=b, oi=oi, residual_cutoff=residual_cutoff ) diff --git a/tests/test_finite.py b/tests/test_finite.py index c6a124d..ac46dba 100644 --- a/tests/test_finite.py +++ b/tests/test_finite.py @@ -7,7 +7,6 @@ """ from collections.abc import Callable -from typing import Any import jax import jax.numpy as jnp @@ -39,8 +38,12 @@ ) def test_validity_let( D: Callable[ - [float | jax.Array | np.ndarray[Any, Any]], - float | jax.Array | np.ndarray[Any, Any], + [ + float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]] + ], + float | jax.Array | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], ], ) -> None: """Finite solver reproduces the reference solution for LET-type models. diff --git a/tests/test_neural/test_exact.py b/tests/test_neural/test_exact.py index 072f09b..5972440 100644 --- a/tests/test_neural/test_exact.py +++ b/tests/test_neural/test_exact.py @@ -8,8 +8,6 @@ We check that the neural fit reproduces the reference curve. """ -from typing import Any - import jax import jax.numpy as jnp import numpy as np @@ -24,7 +22,11 @@ def test_exact() -> None: """The neural fit reproduces the exact exp(-o) solution within tolerance.""" - def D(theta: float | jax.Array | np.ndarray[Any, Any]) -> float | jax.Array: + def D( + theta: float + | jax.Array + | np.ndarray[tuple[int], np.dtype[np.floating | np.integer]], + ) -> float | jax.Array: return (1 - jnp.log(theta)) / 2 o = np.linspace(0, 20, 100)