diff --git a/sbi/inference/posteriors/vi_posterior.py b/sbi/inference/posteriors/vi_posterior.py index d10162722..80e0cf228 100644 --- a/sbi/inference/posteriors/vi_posterior.py +++ b/sbi/inference/posteriors/vi_posterior.py @@ -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. " diff --git a/sbi/inference/potentials/base_potential.py b/sbi/inference/potentials/base_potential.py index f353d535f..7b9676812 100644 --- a/sbi/inference/potentials/base_potential.py +++ b/sbi/inference/potentials/base_potential.py @@ -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 @abstractmethod def __call__(self, theta: Tensor, track_gradients: bool = True) -> Tensor: diff --git a/sbi/inference/potentials/likelihood_based_potential.py b/sbi/inference/potentials/likelihood_based_potential.py index 991254489..ce15d204e 100644 --- a/sbi/inference/potentials/likelihood_based_potential.py +++ b/sbi/inference/potentials/likelihood_based_potential.py @@ -19,6 +19,7 @@ ) from sbi.sbi_types import TorchTransform from sbi.utils.sbiutils import mcmc_transform +from sbi.utils.user_input_checks import process_x def likelihood_estimator_based_potential( @@ -102,7 +103,9 @@ def bind(self, x_o: Tensor, x_is_iid: bool = True) -> "LikelihoodBasedPotential" 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: diff --git a/sbi/inference/potentials/posterior_based_potential.py b/sbi/inference/potentials/posterior_based_potential.py index 91a4af0e3..791bc8c30 100644 --- a/sbi/inference/potentials/posterior_based_potential.py +++ b/sbi/inference/potentials/posterior_based_potential.py @@ -21,6 +21,7 @@ within_support, ) from sbi.utils.torchutils import ensure_theta_batched, infer_module_device +from sbi.utils.user_input_checks import process_x def posterior_estimator_based_potential( @@ -108,14 +109,19 @@ def bind(self, x_o: Tensor, x_is_iid: bool = False) -> "PosteriorBasedPotential" 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: + 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. diff --git a/sbi/inference/potentials/ratio_based_potential.py b/sbi/inference/potentials/ratio_based_potential.py index 917fede47..6b328c933 100644 --- a/sbi/inference/potentials/ratio_based_potential.py +++ b/sbi/inference/potentials/ratio_based_potential.py @@ -11,6 +11,7 @@ from sbi.sbi_types import TorchTransform from sbi.utils.sbiutils import match_theta_and_x_batch_shapes, mcmc_transform from sbi.utils.torchutils import atleast_2d +from sbi.utils.user_input_checks import process_x def ratio_estimator_based_potential( @@ -90,7 +91,9 @@ def bind(self, x_o: Tensor, x_is_iid: bool = True) -> "RatioBasedPotential": 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: diff --git a/sbi/inference/potentials/vector_field_potential.py b/sbi/inference/potentials/vector_field_potential.py index 50681228f..8c492eb53 100644 --- a/sbi/inference/potentials/vector_field_potential.py +++ b/sbi/inference/potentials/vector_field_potential.py @@ -22,6 +22,7 @@ from sbi.sbi_types import TorchTransform from sbi.utils.sbiutils import mcmc_transform, within_support from sbi.utils.torchutils import ensure_theta_batched +from sbi.utils.user_input_checks import process_x class VectorFieldBasedPotential(BasePotential): @@ -61,6 +62,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 {} @@ -119,7 +122,10 @@ 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: + 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 @@ -145,22 +151,24 @@ def bind( prior=self.prior, x_o=None, device=self.device, - iid_method=self.iid_method, - iid_params=self.iid_params, + iid_method=iid_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__( diff --git a/sbi/samplers/vi/vi_divergence_optimizers.py b/sbi/samplers/vi/vi_divergence_optimizers.py index 24ab3d9f1..e57ab621a 100644 --- a/sbi/samplers/vi/vi_divergence_optimizers.py +++ b/sbi/samplers/vi/vi_divergence_optimizers.py @@ -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 @@ -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 diff --git a/sbi/utils/torchutils.py b/sbi/utils/torchutils.py index 4405bed65..ce657cf34 100644 --- a/sbi/utils/torchutils.py +++ b/sbi/utils/torchutils.py @@ -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) @@ -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 @@ -715,7 +716,7 @@ def _base_recursor( key=k, check=check, action=action, - _active=_active, + _active=_active_holder, ) elif isinstance(obj, (List, Tuple, Generator)): new_obj = [] @@ -723,12 +724,12 @@ def _base_recursor( 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: diff --git a/sbi/utils/user_input_checks_utils.py b/sbi/utils/user_input_checks_utils.py index dbd8a1d5e..4a6ff8d0e 100644 --- a/sbi/utils/user_input_checks_utils.py +++ b/sbi/utils/user_input_checks_utils.py @@ -181,6 +181,9 @@ 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, @@ -189,10 +192,6 @@ def __init__( ), ) - 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),