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
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ build-backend = "hatchling.build"

[project]
name = "nlls-gram"
version = "2.7.0"
version = "2.8.0"
description = "Metric-aware Levenberg-Marquardt nonlinear least-squares for JAX"
readme = "README.md"
license = "MIT"
Expand Down
6 changes: 4 additions & 2 deletions src/nlls_gram/multi_start.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,7 +92,9 @@ class DrawNNXModule:

Given a ``MultiStart`` retry key, builds
``module_cls(*args, rngs=nnx.Rngs(key), **kwargs)`` and returns its ``nnx.Param``
state as the new solver start, passing ``args`` through unchanged. Use it instead
state as the new solver start, passing ``args`` through unchanged. Non-``Param``
Variables (scaling constants, statistics) are excluded from the drawn state --
the residual's ``nnx.merge`` supplies them alongside the graphdef. Use it instead
of hand-rolling a re-init closure per driver::

draw = DrawNNXModule(SequentialMLP, settings, dtype=dtype)
Expand Down Expand Up @@ -122,7 +124,7 @@ def __call__(self, key, x_old, args_old):
from flax import nnx

module = self.module_cls(*self.args, rngs=nnx.Rngs(key), **dict(self.kwargs))
_, theta = nnx.split(module, nnx.Param)
_, theta, _ = nnx.split(module, nnx.Param, ...)
return theta, args_old

def __hash__(self):
Expand Down
37 changes: 37 additions & 0 deletions tests/test_nnx_gram_lm.py
Original file line number Diff line number Diff line change
Expand Up @@ -236,6 +236,43 @@ def run(draw):
assert jnp.allclose(got, want)


class ScaledCurveMLP(nnx.Module):
def __init__(self, *, scale=2.0, rngs: nnx.Rngs):
self.scale = nnx.Variable(jnp.asarray(scale))
self.hidden = nnx.Linear(1, 8, rngs=rngs)
self.head = nnx.Linear(8, 1, rngs=rngs)

def __call__(self, x):
return self.scale[...] * self.head(nnx.tanh(self.hidden(x[:, None])))[:, 0]


def test_draw_nnx_module_excludes_non_param_variables():
graphdef, theta_0, nondiff = nnx.split(
ScaledCurveMLP(rngs=nnx.Rngs(1)), nnx.Param, ...
)
assert len(jax.tree.leaves(nondiff)) == 1

draw = DrawNNXModule(ScaledCurveMLP)
theta, args_out = draw(jax.random.key(7), None, ("args",))
assert args_out == ("args",)
# Only the Param leaves are drawn; the residual's merge supplies nondiff.
assert jax.tree_util.tree_structure(theta) == jax.tree_util.tree_structure(theta_0)

ts = jnp.linspace(-1.0, 1.0, 32)
ys = jnp.sin(2.0 * ts)
theta_bad = jax.tree.map(lambda leaf: leaf * jnp.nan, theta_0)

def residual(theta, args, p):
ts, ys = args
return nnx.merge(graphdef, theta, nondiff)(ts) - ys

solver = LevenbergMarquardt(residual, init_damping=1e-2)
ms = MultiStart(key=jax.random.key(2), num_starts=4, draw=draw)
result = solver.solve(theta_bad, (ts, ys), max_steps=200, atol=5e-3, multi_start=ms)
assert int(result.status) == LMStatus.CONVERGED
assert bool(result.multi_start.accepted)


@pytest.mark.parametrize("parallel", [False, True])
def test_multi_start_draw_nnx_module_recovers_from_bad_init(parallel):
ts = jnp.linspace(-1.0, 1.0, 32)
Expand Down
2 changes: 1 addition & 1 deletion uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.