From c4ab381d292e189e41666a71700cc6e7b0551c3f Mon Sep 17 00:00:00 2001 From: Jan Date: Mon, 20 Jul 2026 21:51:20 +0200 Subject: [PATCH 1/2] fix: make unconstraining transforms module-owned --- .../estimators/mixture_density_estimator.py | 17 ++-- sbi/neural_nets/net_builders/flow.py | 4 +- sbi/utils/sbiutils.py | 88 +++++++++++-------- tests/inference_on_device_test.py | 57 +++++++----- tests/save_and_load_test.py | 23 +++++ tests/sbiutils_test.py | 26 ++++++ 6 files changed, 147 insertions(+), 68 deletions(-) diff --git a/sbi/neural_nets/estimators/mixture_density_estimator.py b/sbi/neural_nets/estimators/mixture_density_estimator.py index 2792144b9..1dda5c892 100644 --- a/sbi/neural_nets/estimators/mixture_density_estimator.py +++ b/sbi/neural_nets/estimators/mixture_density_estimator.py @@ -22,7 +22,7 @@ from sbi.neural_nets.estimators.base import ConditionalDensityEstimator from sbi.neural_nets.estimators.mog import MoG from sbi.sbi_types import TorchTransform -from sbi.utils.sbiutils import _apply_to_transform +from sbi.utils.sbiutils import CallableTransform class MultivariateGaussianMDN(nn.Module): @@ -384,7 +384,9 @@ def __init__( f"MDN context_features ({net.context_features})" ) - self._prior_transform = prior_transform + self._prior_transform_module = ( + CallableTransform(prior_transform) if prior_transform is not None else None + ) # Store z-score transform parameters as buffers (not trained, moved with model) if transform_input is not None: @@ -406,11 +408,12 @@ def __init__( self.register_buffer("_transform_shift", None) self.register_buffer("_transform_scale", None) - def _apply(self, fn): - super()._apply(fn) - if self._prior_transform is not None: - _apply_to_transform(self._prior_transform, fn) - return self + @property + def _prior_transform(self) -> Optional[TorchTransform]: + """Return the constrained-to-unconstrained input transform, if configured.""" + if self._prior_transform_module is None: + return None + return self._prior_transform_module.transform @property def embedding_net(self) -> nn.Module: diff --git a/sbi/neural_nets/net_builders/flow.py b/sbi/neural_nets/net_builders/flow.py index e443b6160..46749c7e7 100644 --- a/sbi/neural_nets/net_builders/flow.py +++ b/sbi/neural_nets/net_builders/flow.py @@ -1306,7 +1306,9 @@ def _prepare_x_transforms( "`x_dist` requires a `.support` attribute for" "an unconstrained transformation." ) - transform_to_unconstrained = biject_transform_zuko(mcmc_transform(x_dist)) + transform_to_unconstrained = biject_transform_zuko( + mcmc_transform(x_dist, device=batch_x.device) + ) transforms = (transform_to_unconstrained,) elif z_score_x_bool: z_score_transform = standardizing_transform_zuko(batch_x, structured_x) diff --git a/sbi/utils/sbiutils.py b/sbi/utils/sbiutils.py index 02cdbf6f3..6e6775955 100644 --- a/sbi/utils/sbiutils.py +++ b/sbi/utils/sbiutils.py @@ -286,15 +286,61 @@ def standardizing_transform_zuko( ) -class CallableTransform: - """Wraps a PyTorch Transform to be used in Zuko UnconditionalTransform.""" +def _contains_inverse_transform(transform: TorchTransform) -> bool: + """Return whether a transform tree contains an inverse wrapper.""" + seen = set() - def __init__(self, transform): - self.transform = transform + def contains_inverse(current: TorchTransform) -> bool: + if id(current) in seen: + return False + seen.add(id(current)) + if isinstance(current, torch_tf._InverseTransform): + return True + for value in current.__dict__.values(): + if isinstance(value, TorchTransform) and contains_inverse(value): + return True + if isinstance(value, (list, tuple)) and any( + isinstance(item, TorchTransform) and contains_inverse(item) + for item in value + ): + return True + return False + + return contains_inverse(transform) + + +class CallableTransform(nn.Module): + """Own a PyTorch transform as a movable, picklable module. + + Transform tensors are intentionally absent from ``state_dict``; callers rebuild + the transform before loading estimator weights. + """ - def __call__(self): + def __init__(self, transform: TorchTransform): + super().__init__() + self._is_inverse = _contains_inverse_transform(transform) + self._transform = transform.inv if self._is_inverse else transform + if not isinstance( + self._transform, TorchTransform + ) or _contains_inverse_transform(self._transform): + raise ValueError( + "CallableTransform requires one transform orientation without " + "nested inverse wrappers." + ) + + @property + def transform(self) -> TorchTransform: + """Return the transform in the orientation supplied at construction.""" + return self._transform.inv if self._is_inverse else self._transform + + def forward(self) -> TorchTransform: return self.transform + def _apply(self, fn): + super()._apply(fn) + _apply_to_transform(self._transform, fn) + return self + def biject_transform_zuko( transform: TorchTransform, @@ -842,38 +888,6 @@ def _walk(t): _walk(transform) -def _transform_tensors( - transform: TorchTransform, -) -> List[Tensor]: - """Collect every tensor in a transform tree (mirror of _apply_to_transform). - - Args: - transform: Root of the transform tree. - - Returns: - List of all tensors found in the transform tree. - """ - seen = set() - tensors: List[Tensor] = [] - - def _walk(t): - if id(t) in seen: - return - seen.add(id(t)) - for val in t.__dict__.values(): - if isinstance(val, Tensor): - tensors.append(val) - elif isinstance(val, (list, tuple)): - for item in val: - if isinstance(item, TorchTransform): - _walk(item) - elif isinstance(val, TorchTransform): - _walk(val) - - _walk(transform) - return tensors - - def mcmc_transform( prior: Distribution, num_prior_samples_for_zscoring: int = 1000, diff --git a/tests/inference_on_device_test.py b/tests/inference_on_device_test.py index 395f09a71..311792cdb 100644 --- a/tests/inference_on_device_test.py +++ b/tests/inference_on_device_test.py @@ -58,6 +58,19 @@ ) +def _collect_transform_tensors(transform): + from sbi.utils.sbiutils import _apply_to_transform + + tensors = [] + + def collect(tensor): + tensors.append(tensor) + return tensor + + _apply_to_transform(transform, collect) + return tensors + + @pytest.mark.slow @pytest.mark.gpu @pytest.mark.parametrize( @@ -893,11 +906,9 @@ def test_npe_pfn_on_device(prior_device): ) -@pytest.mark.gpu def test_mdn_device_transform(): """MDN with transform_to_unconstrained moves transform tensors on .to().""" from sbi.neural_nets.net_builders.mdn import build_mdn - from sbi.utils.sbiutils import _transform_tensors device = process_device("gpu") prior = BoxUniform(-2 * torch.ones(2), 2 * torch.ones(2)) @@ -905,7 +916,7 @@ def test_mdn_device_transform(): est = build_mdn(bx, by, z_score_x="transform_to_unconstrained", x_dist=prior) est.to(device) - transform_tensors = _transform_tensors(est._prior_transform) + transform_tensors = _collect_transform_tensors(est._prior_transform) assert transform_tensors, "expected the prior transform to hold tensors" for t in transform_tensors: assert t.device.type == device.split(":")[0], ( @@ -920,27 +931,27 @@ def test_mdn_device_transform(): assert s.device.type == device.split(":")[0] -def test_mdn_transform_follows_dtype(): - """transform_to_unconstrained transform follows dtype casts on the MDN. - - Guards the callable-fn design in _apply_to_transform: a rebuild-from-prior - would leave the transform in float32 and silently desync it from the weights. - Runs on every CI (no GPU needed). - """ - from sbi.neural_nets.net_builders.mdn import build_mdn - from sbi.utils.sbiutils import _transform_tensors +def test_zuko_device_transform(): + """Zuko's unconstraining transform follows accelerator device moves.""" + from sbi.neural_nets.net_builders.flow import build_zuko_maf + from sbi.utils.sbiutils import CallableTransform + device = process_device("gpu") prior = BoxUniform(-2 * torch.ones(2), 2 * torch.ones(2)) - bx, by = prior.sample((256,)), torch.randn(256, 3) - est = build_mdn(bx, by, z_score_x="transform_to_unconstrained", x_dist=prior) - - est.double() + bx, by = prior.sample((512,)), torch.randn(512, 3) + est = build_zuko_maf(bx, by, z_score_x="transform_to_unconstrained", x_dist=prior) + est.to(device) - transform_tensors = _transform_tensors(est._prior_transform) - assert transform_tensors, "expected the prior transform to hold tensors" - assert all(t.dtype == torch.float64 for t in transform_tensors) + wrappers = [ + module for module in est.modules() if isinstance(module, CallableTransform) + ] + assert len(wrappers) == 1 + assert all( + tensor.device.type == device.split(":")[0] + for tensor in _collect_transform_tensors(wrappers[0].transform) + ) - theta = prior.sample((5,)).double() - cond = torch.randn(1, 3).double() - lp = est.log_prob(theta.unsqueeze(1), cond) - assert lp.dtype == torch.float64 + theta = prior.sample((5,)).to(device) + cond = torch.randn(1, 3).to(device) + assert est.log_prob(theta.unsqueeze(1), cond).device.type == device.split(":")[0] + assert est.sample((10,), cond).device.type == device.split(":")[0] diff --git a/tests/save_and_load_test.py b/tests/save_and_load_test.py index 5e26c88c6..c92fbb52b 100644 --- a/tests/save_and_load_test.py +++ b/tests/save_and_load_test.py @@ -68,3 +68,26 @@ def test_picklability( pickle.dump(inference, handle) with open(f"{tmp_path}/saved_inference.pickle", "rb") as handle: _ = pickle.load(handle) + + +@pytest.mark.parametrize("builder_name", ["mdn", "zuko"]) +def test_unconstraining_transform_survives_pickle_and_dtype_cast(builder_name): + """Transformed MDN and Zuko estimators remain usable after serialization.""" + from sbi.neural_nets.net_builders.flow import build_zuko_maf + from sbi.neural_nets.net_builders.mdn import build_mdn + from sbi.utils import BoxUniform + + builder = build_mdn if builder_name == "mdn" else build_zuko_maf + prior = BoxUniform(-2 * torch.ones(2), 2 * torch.ones(2)) + bx, by = prior.sample((256,)), torch.randn(256, 3) + estimator = builder(bx, by, z_score_x="transform_to_unconstrained", x_dist=prior) + + theta, condition = prior.sample((5,)).unsqueeze(1), torch.randn(1, 3) + expected = estimator.log_prob(theta, condition) + reloaded = pickle.loads(pickle.dumps(estimator)) + assert torch.allclose(reloaded.log_prob(theta, condition), expected) + + reloaded.double() + actual = reloaded.log_prob(theta.double(), condition.double()) + assert actual.dtype == torch.float64 + assert torch.isfinite(actual).all() diff --git a/tests/sbiutils_test.py b/tests/sbiutils_test.py index fd8de5634..be384d680 100644 --- a/tests/sbiutils_test.py +++ b/tests/sbiutils_test.py @@ -709,3 +709,29 @@ def test_mdn_transform_to_unconstrained(): assert torch.allclose(lp, mog.log_prob(z) + ldj, atol=1e-5) s = est.sample((10,), cond) assert s.shape[0] == 10 and torch.isfinite(s).all() + + +@pytest.mark.parametrize("orientation", ["concrete", "inverse", "composed_inverse"]) +def test_callable_transform_preserves_orientation_and_dtype(orientation): + """CallableTransform safely owns concrete and inverse transforms.""" + import pickle + + from torch.distributions.transforms import AffineTransform, ComposeTransform + + from sbi.utils.sbiutils import CallableTransform + + transform = AffineTransform(torch.zeros(2), 2 * torch.ones(2)) + if orientation == "inverse": + transform = transform.inv + elif orientation == "composed_inverse": + transform = ComposeTransform([ + transform, + AffineTransform(torch.ones(2), 3 * torch.ones(2)), + ]).inv + + wrapped = pickle.loads(pickle.dumps(CallableTransform(transform))) + x = torch.ones(2) + assert torch.allclose(wrapped.transform(x), transform(x)) + + wrapped.double() + assert wrapped.transform(x.double()).dtype == torch.float64 From 1f95b0f51d3c1f49fcd4de077f5a819584ca4017 Mon Sep 17 00:00:00 2001 From: Jan Date: Tue, 21 Jul 2026 08:29:20 +0200 Subject: [PATCH 2/2] fix: ignore cached inverse transform references --- sbi/utils/sbiutils.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/sbi/utils/sbiutils.py b/sbi/utils/sbiutils.py index 6e6775955..a091332d3 100644 --- a/sbi/utils/sbiutils.py +++ b/sbi/utils/sbiutils.py @@ -296,7 +296,9 @@ def contains_inverse(current: TorchTransform) -> bool: seen.add(id(current)) if isinstance(current, torch_tf._InverseTransform): return True - for value in current.__dict__.values(): + for key, value in current.__dict__.items(): + if key == "_inv": + continue if isinstance(value, TorchTransform) and contains_inverse(value): return True if isinstance(value, (list, tuple)) and any(