Skip to content
Open
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
17 changes: 10 additions & 7 deletions sbi/neural_nets/estimators/mixture_density_estimator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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:
Expand All @@ -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:
Expand Down
4 changes: 3 additions & 1 deletion sbi/neural_nets/net_builders/flow.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
90 changes: 53 additions & 37 deletions sbi/utils/sbiutils.py
Original file line number Diff line number Diff line change
Expand Up @@ -286,15 +286,63 @@ 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 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(
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,
Expand Down Expand Up @@ -842,38 +890,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,
Expand Down
57 changes: 34 additions & 23 deletions tests/inference_on_device_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -893,19 +906,17 @@ 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))
bx, by = prior.sample((512,)), torch.randn(512, 3)
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], (
Expand All @@ -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]
23 changes: 23 additions & 0 deletions tests/save_and_load_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
26 changes: 26 additions & 0 deletions tests/sbiutils_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading