Skip to content

Commit d371c37

Browse files
author
Tobias Buck
committed
feat(parameterization): add constrained transforms for age and metallicity
1 parent 722f854 commit d371c37

4 files changed

Lines changed: 261 additions & 3 deletions

File tree

docs/development/gradient_production_pr_plan.md

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -15,14 +15,14 @@ proof-of-concept notebooks to production full-IFU workflows.
1515

1616
## PR stack
1717

18-
1. `feat(inference-api)` (current)
18+
1. `feat(inference-api)` (completed)
1919
- Add `rubix.inference` module with:
2020
- parameter application without mutating baseline `RubixData`
2121
- `forward`, `loss`, and `value_and_grad` API
2222
- Add unit tests for copy semantics, deterministic loss calls, and analytical
2323
gradient checks.
2424

25-
2. `feat(parameterization)`
25+
2. `feat(parameterization)` (current)
2626
- Add constrained parameter transforms (age/metallicity first) and tests.
2727

2828
3. `refactor(gradient-modes)`

rubix/inference/__init__.py

Lines changed: 22 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,26 @@
11
"""Inference helpers for gradient-based modeling workflows."""
22

33
from .api import apply_params, forward, loss, value_and_grad
4+
from .parameterization import (
5+
IdentityTransform,
6+
ParameterTransform,
7+
SigmoidBounds,
8+
SoftplusLowerBound,
9+
apply_transforms,
10+
build_age_metallicity_transforms,
11+
inverse_transforms,
12+
)
413

5-
__all__ = ["apply_params", "forward", "loss", "value_and_grad"]
14+
__all__ = [
15+
"IdentityTransform",
16+
"ParameterTransform",
17+
"SigmoidBounds",
18+
"SoftplusLowerBound",
19+
"apply_params",
20+
"apply_transforms",
21+
"build_age_metallicity_transforms",
22+
"forward",
23+
"inverse_transforms",
24+
"loss",
25+
"value_and_grad",
26+
]
Lines changed: 148 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,148 @@
1+
from dataclasses import dataclass
2+
from typing import Mapping
3+
4+
import jax
5+
import jax.numpy as jnp
6+
from beartype.typing import Any
7+
8+
ParamsTree = Mapping[str, Mapping[str, Any]]
9+
TransformTree = Mapping[str, Mapping[str, "ParameterTransform"]]
10+
11+
12+
class ParameterTransform:
13+
"""Base class for constrained parameter transforms."""
14+
15+
def forward(self, unconstrained: Any) -> Any:
16+
"""Map unconstrained values to constrained space."""
17+
raise NotImplementedError
18+
19+
def inverse(self, constrained: Any) -> Any:
20+
"""Map constrained values to unconstrained space."""
21+
raise NotImplementedError
22+
23+
24+
@dataclass(frozen=True)
25+
class IdentityTransform(ParameterTransform):
26+
"""No-op transform."""
27+
28+
def forward(self, unconstrained: Any) -> Any:
29+
return unconstrained
30+
31+
def inverse(self, constrained: Any) -> Any:
32+
return constrained
33+
34+
35+
@dataclass(frozen=True)
36+
class SoftplusLowerBound(ParameterTransform):
37+
"""Transform values to ``(lower, inf)`` via softplus."""
38+
39+
lower: float
40+
eps: float = 1e-8
41+
42+
def forward(self, unconstrained: Any) -> Any:
43+
return self.lower + jax.nn.softplus(unconstrained) + self.eps
44+
45+
def inverse(self, constrained: Any) -> Any:
46+
shifted = constrained - self.lower - self.eps
47+
safe = jnp.maximum(shifted, jnp.asarray(self.eps, dtype=shifted.dtype))
48+
return jnp.log(jnp.expm1(safe))
49+
50+
51+
@dataclass(frozen=True)
52+
class SigmoidBounds(ParameterTransform):
53+
"""Transform values to ``(lower, upper)`` via sigmoid."""
54+
55+
lower: float
56+
upper: float
57+
eps: float = 1e-8
58+
59+
def __post_init__(self):
60+
if self.upper <= self.lower:
61+
raise ValueError("upper must be strictly larger than lower")
62+
63+
def forward(self, unconstrained: Any) -> Any:
64+
width = self.upper - self.lower
65+
return self.lower + width * jax.nn.sigmoid(unconstrained)
66+
67+
def inverse(self, constrained: Any) -> Any:
68+
width = self.upper - self.lower
69+
x = (constrained - self.lower) / width
70+
x = jnp.clip(x, self.eps, 1.0 - self.eps)
71+
return jnp.log(x) - jnp.log1p(-x)
72+
73+
74+
def apply_transforms(
75+
params: ParamsTree,
76+
transforms: TransformTree,
77+
direction: str = "forward",
78+
) -> dict[str, dict[str, Any]]:
79+
"""Apply transform tree to a parameter tree.
80+
81+
Args:
82+
params (ParamsTree): Nested parameter dictionary.
83+
transforms (TransformTree): Nested transform dictionary with the same
84+
keys as ``params`` for transformed leaves.
85+
direction (str, optional): Either ``"forward"`` (unconstrained ->
86+
constrained) or ``"inverse"`` (constrained -> unconstrained).
87+
Defaults to ``"forward"``.
88+
89+
Raises:
90+
ValueError: If ``direction`` is invalid.
91+
92+
Returns:
93+
dict[str, dict[str, Any]]: Transformed parameter dictionary.
94+
"""
95+
if direction not in {"forward", "inverse"}:
96+
raise ValueError("direction must be one of {'forward', 'inverse'}")
97+
98+
transformed: dict[str, dict[str, Any]] = {}
99+
for component, fields in params.items():
100+
transformed[component] = {}
101+
component_transforms = transforms.get(component, {})
102+
for field, value in fields.items():
103+
transform = component_transforms.get(field, IdentityTransform())
104+
if direction == "forward":
105+
transformed[component][field] = transform.forward(value)
106+
else:
107+
transformed[component][field] = transform.inverse(value)
108+
109+
return transformed
110+
111+
112+
def inverse_transforms(
113+
params: ParamsTree,
114+
transforms: TransformTree,
115+
) -> dict[str, dict[str, Any]]:
116+
"""Apply inverse transforms to a parameter tree."""
117+
return apply_transforms(params=params, transforms=transforms, direction="inverse")
118+
119+
120+
def build_age_metallicity_transforms(
121+
age_lower: float = 0.0,
122+
age_upper: float = 20.0,
123+
metallicity_lower: float = 0.0,
124+
metallicity_upper: float = 0.05,
125+
) -> dict[str, dict[str, ParameterTransform]]:
126+
"""Build default transform tree for age and metallicity optimization.
127+
128+
Args:
129+
age_lower (float, optional): Lower age bound in Gyr. Defaults to 0.0.
130+
age_upper (float, optional): Upper age bound in Gyr. Defaults to 20.0.
131+
metallicity_lower (float, optional): Lower metallicity bound.
132+
Defaults to 0.0.
133+
metallicity_upper (float, optional): Upper metallicity bound.
134+
Defaults to 0.05.
135+
136+
Returns:
137+
dict[str, dict[str, ParameterTransform]]: Transform tree for
138+
``stars.age`` and ``stars.metallicity``.
139+
"""
140+
return {
141+
"stars": {
142+
"age": SigmoidBounds(lower=age_lower, upper=age_upper),
143+
"metallicity": SigmoidBounds(
144+
lower=metallicity_lower,
145+
upper=metallicity_upper,
146+
),
147+
}
148+
}
Lines changed: 89 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,89 @@
1+
import jax.numpy as jnp
2+
import pytest
3+
4+
from rubix.inference import (
5+
IdentityTransform,
6+
SigmoidBounds,
7+
SoftplusLowerBound,
8+
apply_transforms,
9+
build_age_metallicity_transforms,
10+
inverse_transforms,
11+
)
12+
13+
14+
def test_identity_transform_roundtrip():
15+
transform = IdentityTransform()
16+
values = jnp.array([-2.0, 0.0, 3.0])
17+
18+
constrained = transform.forward(values)
19+
recovered = transform.inverse(constrained)
20+
21+
assert jnp.allclose(constrained, values)
22+
assert jnp.allclose(recovered, values)
23+
24+
25+
def test_sigmoid_bounds_forward_and_inverse():
26+
transform = SigmoidBounds(lower=0.0, upper=20.0)
27+
unconstrained = jnp.array([-3.0, 0.0, 3.0])
28+
29+
constrained = transform.forward(unconstrained)
30+
recovered = transform.inverse(constrained)
31+
32+
assert jnp.all(constrained > 0.0)
33+
assert jnp.all(constrained < 20.0)
34+
assert jnp.allclose(recovered, unconstrained, atol=1e-5, rtol=1e-5)
35+
36+
37+
def test_softplus_lower_bound_forward_and_inverse():
38+
transform = SoftplusLowerBound(lower=0.0)
39+
unconstrained = jnp.array([-6.0, 0.0, 4.0])
40+
41+
constrained = transform.forward(unconstrained)
42+
recovered = transform.inverse(constrained)
43+
44+
assert jnp.all(constrained > 0.0)
45+
assert jnp.allclose(recovered, unconstrained, atol=1e-5, rtol=1e-5)
46+
47+
48+
def test_apply_transforms_tree_roundtrip():
49+
params = {
50+
"stars": {
51+
"age": jnp.array([-1.0, 0.0, 1.0]),
52+
"metallicity": jnp.array([-2.0, 0.5, 2.0]),
53+
"mass": jnp.array([1.0, 2.0, 3.0]),
54+
}
55+
}
56+
transforms = build_age_metallicity_transforms(
57+
age_lower=0.0,
58+
age_upper=20.0,
59+
metallicity_lower=0.0,
60+
metallicity_upper=0.05,
61+
)
62+
63+
constrained = apply_transforms(params, transforms, direction="forward")
64+
recovered = inverse_transforms(constrained, transforms)
65+
66+
assert jnp.all(constrained["stars"]["age"] > 0.0)
67+
assert jnp.all(constrained["stars"]["age"] < 20.0)
68+
assert jnp.all(constrained["stars"]["metallicity"] > 0.0)
69+
assert jnp.all(constrained["stars"]["metallicity"] < 0.05)
70+
assert jnp.allclose(recovered["stars"]["age"], params["stars"]["age"], atol=1e-5)
71+
assert jnp.allclose(
72+
recovered["stars"]["metallicity"],
73+
params["stars"]["metallicity"],
74+
atol=1e-5,
75+
)
76+
assert jnp.allclose(recovered["stars"]["mass"], params["stars"]["mass"])
77+
78+
79+
def test_apply_transforms_raises_on_bad_direction():
80+
params = {"stars": {"age": jnp.array([0.0])}}
81+
transforms = build_age_metallicity_transforms()
82+
83+
with pytest.raises(ValueError, match="direction must be one of"):
84+
apply_transforms(params, transforms, direction="bad")
85+
86+
87+
def test_sigmoid_bounds_raises_for_invalid_bounds():
88+
with pytest.raises(ValueError, match="upper must be strictly larger than lower"):
89+
SigmoidBounds(lower=1.0, upper=1.0)

0 commit comments

Comments
 (0)