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
1 change: 1 addition & 0 deletions .pre-commit-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ repos:
hooks:
- id: check-ast
- id: check-added-large-files
args: ['--maxkb=10000'] # 10MB limit — notebooks with plots are commonly several MB
- id: check-merge-conflict
- id: check-case-conflict
- id: check-yaml
Expand Down
3 changes: 3 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@ The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
### Added

- Added comprehensive notebook development guidelines in `docs/notebooks/README.md`
- get_implementations() and has_implementation() for operator dispatch introspection
- Operator params guide tutorial covering inspection, override, shared params, eqx.Module as params, NN-generated stencils, and learned gradient correction

### Changed

Expand All @@ -32,6 +34,7 @@ The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),

- Fixed bug in `Continuous.replace_params` where deprecated `self.get_field` was used instead of `self.get_fun`
- Inconsistent shift_operator between FourierSeries and FiniteDifferences (#146)
- Linear.__eq__ no longer crashes when compared with non-Linear types (fixes #145)

## [0.2.8] - 2024-09-17

Expand Down
1,021 changes: 1,021 additions & 0 deletions docs/notebooks/operator_params_guide.ipynb

Large diffs are not rendered by default.

2 changes: 1 addition & 1 deletion docs/util.md
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
handler: python
selection:
filters:
- "__init__$"
- "!^_"
rendering:
show_root_heading: true
show_source: false
Expand Down
3 changes: 3 additions & 0 deletions jaxdf/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
from jaxdf.core import Field # isort:skip

from jaxdf import util, geometry, mods, operators # isort:skip
from jaxdf.util import get_implementations, has_implementation # isort:skip

# Must be imported after discretization
from jaxdf.operators.magic import * # isort:skip
Expand All @@ -35,4 +36,6 @@
'Linear',
'Module',
'OnGrid',
'get_implementations',
'has_implementation',
]
2 changes: 2 additions & 0 deletions jaxdf/discretization.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,8 @@ class Linear(Field):
domain: Domain

def __eq__(self, other):
if not isinstance(other, Linear):
return False
return tree_equal(self, other) * (self.domain == other.domain)

@property
Expand Down
44 changes: 44 additions & 0 deletions jaxdf/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,3 +49,47 @@ def get_implemented(f):
instances = set(instances)
for instance in instances:
print(" ─ " + instance)


def get_implementations(f):
"""Returns the implemented type signatures for an operator.

Args:
f: An operator function registered via @operator.

Returns:
list[tuple[str, ...]]: List of type signature tuples for each
implementation.

Example:
>>> from jaxdf.operators import gradient
>>> get_implementations(gradient)
[('Continuous',), ('FiniteDifferences',), ('FourierSeries',)]
"""
instances = []
for f_instance in f.methods:
types = f_instance.signature.types
type_names = tuple(t.__name__ for t in types)
if type_names not in instances:
instances.append(type_names)
return sorted(instances)


def has_implementation(f, *types):
"""Check if an operator has an implementation for the given types.

Args:
f: An operator function registered via @operator.
*types: The types to check for.

Returns:
bool: True if an implementation exists for the given types.

Example:
>>> from jaxdf.operators import gradient
>>> from jaxdf.discretization import FourierSeries
>>> has_implementation(gradient, FourierSeries)
True
"""
type_names = tuple(t.__name__ for t in types)
return type_names in get_implementations(f)
1 change: 1 addition & 0 deletions mkdocs.yml
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,7 @@ nav:
- Home: index.md
- Tutorials:
- Quick start: notebooks/quickstart.ipynb
- Operator params guide: notebooks/operator_params_guide.ipynb
- Physics informed neural networks: notebooks/pinn_burgers.ipynb
- Optimize acoustic simulations: notebooks/simulate_helmholtz_equation.ipynb
- How discretizations work: notebooks/api_discretization.ipynb
Expand Down
8 changes: 8 additions & 0 deletions tests/test_linear.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,3 +33,11 @@ def test_equality():

d = Linear(jnp.asarray([1.0]), Domain((2, ), (1.0, )))
assert a != d


def test_equality_with_type():
"""Regression test for jaxdf#145 — __eq__ should not crash when compared with a type."""
domain = Domain((1, ), (1.0, ))
a = Linear(jnp.asarray([1.0]), domain)
assert (a == Linear) == False
assert (a == int) == False
126 changes: 126 additions & 0 deletions tests/test_module_params.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,126 @@
"""Tests that eqx.Module objects work as operator params with jit/grad/vmap."""
import equinox as eqx
import jax
import jax.numpy as jnp
import pytest

from jaxdf import OnGrid, operator
from jaxdf.geometry import Domain


class SimpleModule(eqx.Module):
weight: jax.Array
bias: jax.Array

def __call__(self, x):
return self.weight * x + self.bias


def module_init(x: OnGrid, *args, **kwargs):
return SimpleModule(weight=jnp.array(2.0), bias=jnp.array(1.0))


@operator(init_params=module_init)
def module_op(x: OnGrid, *, params=None):
"""Operator that uses an eqx.Module as params."""
return x.replace_params(params(x.params))


@pytest.fixture
def field():
domain = Domain((8, ), (1.0, ))
return OnGrid(jnp.ones((8, 1)) * 3.0, domain)


def test_module_params_basic(field):
"""eqx.Module as operator params works."""
result = module_op(field)
expected = 2.0 * 3.0 + 1.0 # weight * value + bias
assert jnp.allclose(result.params, expected)


def test_module_params_explicit(field):
"""Explicitly passing eqx.Module as params works."""
custom = SimpleModule(weight=jnp.array(5.0), bias=jnp.array(0.0))
result = module_op(field, params=custom)
assert jnp.allclose(result.params, 15.0)


def test_module_params_default_params(field):
"""default_params returns the eqx.Module."""
params = module_op.default_params(field)
assert isinstance(params, SimpleModule)
assert params.weight == 2.0


def test_module_params_jit(field):
"""eqx.Module params work with jax.jit."""

@jax.jit
def f(x):
return module_op(x)

result = f(field)
expected = 2.0 * 3.0 + 1.0
assert jnp.allclose(result.params, expected)


def test_module_params_jit_explicit(field):
"""Explicitly passed eqx.Module params work with jax.jit."""
custom = SimpleModule(weight=jnp.array(5.0), bias=jnp.array(0.0))

@jax.jit
def f(x, params):
return module_op(x, params=params)

result = f(field, custom)
assert jnp.allclose(result.params, 15.0)


def test_module_params_grad(field):
"""jax.grad flows through eqx.Module operator params."""

def loss(module_params):
result = module_op(field, params=module_params)
return jnp.sum(result.params)

custom = SimpleModule(weight=jnp.array(3.0), bias=jnp.array(0.0))
grads = jax.grad(loss)(custom)

# d/d(weight) of sum(weight * x + bias) = sum(x)
assert isinstance(grads, SimpleModule)
assert jnp.allclose(grads.weight, jnp.sum(field.params))
# d/d(bias) = n_elements
assert jnp.allclose(grads.bias, float(field.params.size))


def test_module_params_vmap(field):
"""jax.vmap works with eqx.Module operator params (batched over params)."""
weights = jnp.array([1.0, 2.0, 3.0])
biases = jnp.array([0.0, 0.0, 0.0])
batched_modules = jax.vmap(SimpleModule)(weights, biases)

def apply_one(params):
return module_op(field, params=params).params

results = jax.vmap(apply_one)(batched_modules)
assert results.shape == (3, 8, 1)
assert jnp.allclose(results[0], 3.0) # 1.0 * 3.0
assert jnp.allclose(results[1], 6.0) # 2.0 * 3.0
assert jnp.allclose(results[2], 9.0) # 3.0 * 3.0


def test_module_params_jit_grad(field):
"""jit + grad combined works with eqx.Module params."""

@jax.jit
def loss(module_params):
result = module_op(field, params=module_params)
return jnp.sum(result.params)

custom = SimpleModule(weight=jnp.array(3.0), bias=jnp.array(0.0))
grads = jax.grad(loss)(custom)

assert isinstance(grads, SimpleModule)
assert jnp.isfinite(grads.weight)
assert jnp.isfinite(grads.bias)
17 changes: 17 additions & 0 deletions tests/test_util.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,9 @@
from jax import numpy as jnp

from jaxdf import util
from jaxdf.discretization import Continuous, FiniteDifferences, FourierSeries
from jaxdf.operators.differential import gradient
from jaxdf.util import get_implementations, has_implementation


def test_append_dimension():
Expand All @@ -21,5 +23,20 @@ def test_get_implemented():
util.get_implemented(gradient)


def test_get_implementations():
impls = get_implementations(gradient)
assert isinstance(impls, list)
assert len(impls) >= 3 # At least Continuous, FD, Fourier
assert ('FourierSeries', ) in impls
assert ('FiniteDifferences', ) in impls
assert ('Continuous', ) in impls


def test_has_implementation():
assert has_implementation(gradient, FourierSeries) == True
assert has_implementation(gradient, FiniteDifferences) == True
assert has_implementation(gradient, Continuous) == True


if __name__ == "__main__":
test_get_implemented()