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
142 changes: 100 additions & 42 deletions src/frontx/_boltzmann.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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:
Expand All @@ -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
Expand All @@ -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]
70 changes: 44 additions & 26 deletions src/frontx/_forward.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
from collections.abc import Callable
from typing import Any
from typing import cast

import diffrax
import equinox as eqx
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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]
Loading