Skip to content
Merged
Show file tree
Hide file tree
Changes from 11 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
4 changes: 4 additions & 0 deletions sbi/inference/posteriors/vi_posterior.py
Original file line number Diff line number Diff line change
Expand Up @@ -880,6 +880,10 @@ def train(
break
# Training finished:
self._trained_on = x

# Bind potential_fn to the trained x so it can be used for sampling/evaluation
self.potential_fn = self.potential_fn.bind(x)

if self._mode == "amortized":
warnings.warn(
"Switching from amortized to single-x mode. "
Expand Down
5 changes: 4 additions & 1 deletion sbi/inference/potentials/base_potential.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,10 @@ def __init__(
"""
self.device = device
self.prior = prior
self.set_x(x_o)
if x_o is not None:
x_o = process_x(x_o).to(self.device)
self._x_o = x_o
self._x_is_iid = True

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

compared with set_x() we hard to _x_is_iid to True here. Since it's not a method argument, maybe True is a good default. We have to nesure though that it is set correctly in each inherited class. Because right now by calling super() in e.g. the likelihood potential, we will ahve x_is_iid=True by default

@Jocho-Smith Jocho-Smith Jul 23, 2026

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

right now by calling super() in e.g. the likelihood potential, we will ahve x_is_iid=True by default

This is exactly what is happening right now on the main branch. likelihood calls super without _x_is_iid : https://github.com/sbi-dev/sbi/blob/main/sbi/inference/potentials/likelihood_based_potential.py#L80

BasePotential.init always calls set_x without any iid information: https://github.com/sbi-dev/sbi/blob/main/sbi/inference/potentials/base_potential.py#L33

x_is_iid is set with default value True: https://github.com/sbi-dev/sbi/blob/main/sbi/inference/potentials/base_potential.py#L60


@abstractmethod
def __call__(self, theta: Tensor, track_gradients: bool = True) -> Tensor:
Expand Down
6 changes: 5 additions & 1 deletion sbi/inference/potentials/likelihood_based_potential.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,13 +96,17 @@ def to(self, device: Union[str, torch.device]) -> "LikelihoodBasedPotential":

def bind(self, x_o: Tensor, x_is_iid: bool = True) -> "LikelihoodBasedPotential":
"""Create new potential with x bound, without mutable state."""
from sbi.utils.user_input_checks import process_x

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

we should import that not within the function but at file level


bound = LikelihoodBasedPotential(
likelihood_estimator=self.likelihood_estimator,
prior=self.prior,
x_o=None,
device=self.device,
)
bound.set_x(x_o, x_is_iid=x_is_iid)
x_o = process_x(x_o).to(self.device)
bound._x_o = x_o
bound._x_is_iid = x_is_iid
return bound

def __call__(self, theta: Tensor, track_gradients: bool = True) -> Tensor:
Expand Down
13 changes: 11 additions & 2 deletions sbi/inference/potentials/posterior_based_potential.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,20 +102,29 @@ def to(self, device: Union[str, torch.device]) -> "PosteriorBasedPotential":

def bind(self, x_o: Tensor, x_is_iid: bool = False) -> "PosteriorBasedPotential":
"""Create new potential with x bound, without mutable state."""
from sbi.utils.user_input_checks import process_x

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

same is in the likelihoood based potential


bound = PosteriorBasedPotential(
posterior_estimator=self.posterior_estimator,
prior=self.prior,
x_o=None,
device=self.device,
)
bound.set_x(x_o, x_is_iid=x_is_iid)
x_o = process_x(x_o).to(self.device)
bound._x_o = x_o
bound._x_is_iid = x_is_iid
return bound

def set_x(self, x_o: Optional[Tensor], x_is_iid: Optional[bool] = False):
"""
Check the shape of the observed data and, if valid, set it.
"""
super().set_x(x_o, x_is_iid=x_is_iid)
if x_o is not None:
from sbi.utils.user_input_checks import process_x

x_o = process_x(x_o).to(self.device)
self._x_o = x_o
self._x_is_iid = x_is_iid

def __call__(self, theta: Tensor, track_gradients: bool = True) -> Tensor:
r"""Returns the potential for posterior-based methods.
Expand Down
6 changes: 5 additions & 1 deletion sbi/inference/potentials/ratio_based_potential.py
Original file line number Diff line number Diff line change
Expand Up @@ -84,13 +84,17 @@ def to(self, device: Union[str, torch.device]) -> "RatioBasedPotential":

def bind(self, x_o: Tensor, x_is_iid: bool = True) -> "RatioBasedPotential":
"""Create new potential with x bound, without mutable state."""
from sbi.utils.user_input_checks import process_x

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

same is in the likelihoood based potential


bound = RatioBasedPotential(
ratio_estimator=self.ratio_estimator,
prior=self.prior,
x_o=None,
device=self.device,
)
bound.set_x(x_o, x_is_iid=x_is_iid)
x_o = process_x(x_o).to(self.device)
bound._x_o = x_o
bound._x_is_iid = x_is_iid
return bound

def __call__(self, theta: Tensor, track_gradients: bool = True) -> Tensor:
Expand Down
37 changes: 24 additions & 13 deletions sbi/inference/potentials/vector_field_potential.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,8 @@ def __init__(
self.vector_field_estimator.eval()
self.iid_method = iid_method
self.iid_params = iid_params
self.guidance_method: Optional[str] = None
self.guidance_params: Optional[Dict[str, Any]] = None
self.neural_ode_backend = neural_ode_backend

neural_ode_kwargs = neural_ode_kwargs or {}
Expand Down Expand Up @@ -119,7 +121,12 @@ def set_x(
`IIDScoreFunction`.
ode_kwargs: Additional keyword arguments for the neural ODE.
"""
super().set_x(x_o, x_is_iid)
if x_o is not None:
from sbi.utils.user_input_checks import process_x

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Import on file level


x_o = process_x(x_o).to(self.device)
self._x_o = x_o
self._x_is_iid = x_is_iid
self.iid_method = iid_method or self.iid_method
self.iid_params = iid_params
self.guidance_method = guidance_method
Expand All @@ -140,27 +147,31 @@ def bind(
**ode_kwargs,
) -> "VectorFieldBasedPotential":
"""Create new potential with x bound, without mutable state."""
from sbi.utils.user_input_checks import process_x

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

same as above


bound = VectorFieldBasedPotential(
vector_field_estimator=self.vector_field_estimator,
prior=self.prior,
x_o=None,
device=self.device,
iid_method=self.iid_method,
iid_params=self.iid_params,
iid_method=iid_method or self.iid_method,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

this should proabbly also be a id_method if iid_method is not None else self.iid_method,

iid_params=iid_params if iid_params is not None else self.iid_params,
neural_ode_backend=self.neural_ode_backend,
)
bound.neural_ode.params.update(self.neural_ode.params)
# Transfer learned neural ODE parameters (e.g., from training) to the bound
# potential. This preserves the trained flow model when conditioning on new x.
bound.set_x(
x_o,
x_is_iid=x_is_iid,
iid_method=iid_method,
iid_params=iid_params,
guidance_method=guidance_method,
guidance_params=guidance_params,
**ode_kwargs,
x_o = process_x(x_o).to(self.device)
bound._x_o = x_o
bound._x_is_iid = x_is_iid
bound.guidance_method = (
guidance_method if guidance_method is not None else self.guidance_method
)
bound.guidance_params = (
guidance_params if guidance_params is not None else self.guidance_params
)
if not x_is_iid and (bound._x_o is not None):
bound.flow = bound.rebuild_flow(**ode_kwargs)
elif bound._x_o is not None:
bound.flows = bound.rebuild_flows_for_batch(**ode_kwargs)
return bound

def __call__(
Expand Down
4 changes: 2 additions & 2 deletions sbi/samplers/vi/vi_divergence_optimizers.py
Original file line number Diff line number Diff line change
Expand Up @@ -479,7 +479,7 @@ def generate_elbo_particles(
samples = self.q.rsample(torch.Size((num_samples,)))
log_q = self.q.log_prob(samples)

self.potential_fn.set_x(x_o)
self.potential_fn = self.potential_fn.bind(x_o)
log_potential = self.potential_fn(samples)
elbo = log_potential - log_q
return elbo
Expand Down Expand Up @@ -638,7 +638,7 @@ def _loss_q_proposal(self, x_o: Tensor) -> Tuple[Tensor, Tensor]:
if hasattr(self.q, "clear_cache"):
self.q.clear_cache()
logq = self.q.log_prob(samples)
self.potential_fn.set_x(x_o)
self.potential_fn = self.potential_fn.bind(x_o)
logp = self.potential_fn(samples)
with torch.no_grad():
logweights = logp - logq
Expand Down
17 changes: 9 additions & 8 deletions sbi/utils/torchutils.py
Original file line number Diff line number Diff line change
Expand Up @@ -677,13 +677,14 @@ def _base_recursor(
_active: Object identities on the active recursion path. Used internally
to avoid infinite recursion in cyclic object graphs.
"""
if _active is None:
_active = set()
_active_holder: Optional[set[int]] = _active
if _active_holder is None:
_active_holder = set()

obj_id = id(obj)
if obj_id in _active:
if obj_id in _active_holder:
return
_active.add(obj_id)
_active_holder.add(obj_id)

if isinstance(obj, Module) and check(obj):
action(obj)
Expand All @@ -698,7 +699,7 @@ def _base_recursor(
key=k,
check=check,
action=action,
_active=_active,
_active=_active_holder,
)
elif isinstance(obj, type):
# Skip class/type objects to avoid modifying immutable C extension types
Expand All @@ -715,20 +716,20 @@ def _base_recursor(
key=k,
check=check,
action=action,
_active=_active,
_active=_active_holder,
)
elif isinstance(obj, (List, Tuple, Generator)):
new_obj = []
for o in obj:
if check(o):
new_obj.append(action(o))
else:
_base_recursor(o, check=check, action=action, _active=_active)
_base_recursor(o, check=check, action=action, _active=_active_holder)
new_obj.append(o)
if parent is not None and key is not None:
setattr(parent, key, type(obj)(new_obj)) # type: ignore

_active.remove(obj_id)
_active_holder.remove(obj_id)


def move_all_tensor_to_device(obj: object, device: Union[str, torch.device]) -> None:
Expand Down
47 changes: 35 additions & 12 deletions sbi/utils/user_input_checks_utils.py

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Here we seem to have the same modifications that we dicussed in #1943 before.

  • there is code duplication for __deepcopy__
  • I don't see where deepcopy is used and why it is necessary to introduce
  • we silently change the validate_args hard coded to False which seems to be a bug.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

ouh sorry. I messed up the merge. Will fix it.

Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
# under the Apache License Version 2.0, see <https://www.apache.org/licenses/>

import warnings
from copy import deepcopy
from typing import Dict, Optional, Sequence, Union

import torch
Expand Down Expand Up @@ -181,18 +182,15 @@ def __init__(
event_shape=torch.Size(),
validate_args=None,
):
self.prior = prior
self.device = None
self.return_type = return_type
super().__init__(
batch_shape=batch_shape,
event_shape=event_shape,
validate_args=(
prior._validate_args if validate_args is None else validate_args
),
validate_args=False,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

flag (see file level comment)

)

self.prior = prior
self.device = None
self.return_type = return_type

def log_prob(self, value) -> Tensor:
return torch.as_tensor(
self.prior.log_prob(value),
Expand Down Expand Up @@ -236,6 +234,15 @@ def to(self, device: Union[str, torch.device]) -> None:
self.prior = move_distribution_to_device(self.prior, device)
self.device = device

def __deepcopy__(self, memo):
"""Ensure prior attribute is preserved during deepcopy."""
cls = self.__class__
result = cls.__new__(cls)
memo[id(self)] = result
for k, v in self.__dict__.items():
setattr(result, k, deepcopy(v, memo))
return result

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

flag (see file level comment)


class MultipleIndependent(Distribution):
"""Wrap a sequence of PyTorch distributions into a joint PyTorch distribution."""
Expand Down Expand Up @@ -434,6 +441,15 @@ def to(self, device: Union[str, torch.device]) -> None:
self.dists[i] = move_distribution_to_device(self.dists[i], device)
self.device = device

def __deepcopy__(self, memo):
"""Ensure dists attribute is preserved during deepcopy."""
cls = self.__class__
result = cls.__new__(cls)
memo[id(self)] = result
for k, v in self.__dict__.items():
setattr(result, k, deepcopy(v, memo))
return result


def build_support(
lower_bound: Optional[Tensor] = None, upper_bound: Optional[Tensor] = None
Expand Down Expand Up @@ -523,15 +539,13 @@ class OneDimPriorWrapper(Distribution):
"""

def __init__(self, prior: Distribution, validate_args=None) -> None:
self.prior = prior
self.device = None
super().__init__(
batch_shape=prior.batch_shape,
event_shape=prior.event_shape,
validate_args=(
prior._validate_args if validate_args is None else validate_args
),
validate_args=False,
)
self.prior = prior
self.device = None

def to(self, device: Union[str, torch.device]) -> None:
"""
Expand All @@ -546,6 +560,15 @@ def to(self, device: Union[str, torch.device]) -> None:
self.prior = move_distribution_to_device(self.prior, device)
self.device = device

def __deepcopy__(self, memo):
"""Ensure prior attribute is preserved during deepcopy."""
cls = self.__class__
result = cls.__new__(cls)
memo[id(self)] = result
for k, v in self.__dict__.items():
setattr(result, k, deepcopy(v, memo))
return result

def sample(self, *args, **kwargs) -> Tensor:
return self.prior.sample(*args, **kwargs)

Expand Down
Loading