diff --git a/sbi/inference/potentials/base_potential.py b/sbi/inference/potentials/base_potential.py index d24b05b94..8a52dd470 100644 --- a/sbi/inference/potentials/base_potential.py +++ b/sbi/inference/potentials/base_potential.py @@ -69,11 +69,6 @@ def x_o(self) -> Tensor: "No observed data is available. Use `potential_fn.set_x(x_o)`." ) - @x_o.setter - def x_o(self, x_o: Optional[Tensor]) -> None: - """Check the shape of the observed data and, if valid, set it.""" - self.set_x(x_o) - def return_x_o(self) -> Optional[Tensor]: """Return the observed data at which the potential is evaluated. diff --git a/sbi/samplers/vi/vi_divergence_optimizers.py b/sbi/samplers/vi/vi_divergence_optimizers.py index d05262f11..24ab3d9f1 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.x_o = x_o + self.potential_fn.set_x(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.x_o = x_o + self.potential_fn.set_x(x_o) logp = self.potential_fn(samples) with torch.no_grad(): logweights = logp - logq diff --git a/sbi/utils/conditional_density_utils.py b/sbi/utils/conditional_density_utils.py index 0d48efedb..3251f99a9 100644 --- a/sbi/utils/conditional_density_utils.py +++ b/sbi/utils/conditional_density_utils.py @@ -430,11 +430,6 @@ def x_o(self) -> Tensor: else: raise ValueError("No observed data is available.") - @x_o.setter - def x_o(self, x_o: Optional[Tensor]) -> None: - """Check the shape of the observed data and, if valid, set it.""" - self.set_x(x_o) - def return_x_o(self) -> Optional[Tensor]: """Return the observed data at which the potential is evaluated.