Loading src/borch/distributions/distribution_utils.py +13 −3 Original line number Diff line number Diff line Loading @@ -192,6 +192,12 @@ def _verify_is_continuous(rv): if not is_continuous(rv.support): raise RuntimeError(f"{rv} is not continuous") def _dist_mean(dist): try: mean = dist.mean.detach() except NotImplementedError: mean = dist.sample(sample_shape=torch.Size([1000])).mean(0) return mean def normal_distribution_from_rv( rv, log_scale: Union[Tensor, Number], loc_at_mean=True Loading @@ -217,7 +223,7 @@ def normal_distribution_from_rv( """ _verify_is_continuous(rv) if loc_at_mean: loc_init = rv.distribution.mean.detach() loc_init = _dist_mean(rv.distribution).detach() else: loc_init = rv.detach() if not torch.isfinite(loc_init).all(): Loading Loading @@ -263,7 +269,11 @@ def delta_distribution_from_rv(rv): def _scale(rv): if isinstance(rv, (dist.StudentT, dist.Cauchy)): return rv.scale return rv.distribution.stddev try: scale =rv.distribution.stddev except NotImplementedError: scale = rv.distribution.sample(sample_shape=torch.Size([1000])).std() return scale def scaled_normal_dist_from_rv(rv, scaling, loc_at_mean=True): Loading Loading @@ -291,7 +301,7 @@ def scaled_normal_dist_from_rv(rv, scaling, loc_at_mean=True): """ _verify_is_continuous(rv) if loc_at_mean: loc_init = rv.distribution.mean.detach() loc_init = _dist_mean(rv.distribution).detach() else: loc_init = rv.detach() if not torch.isfinite(loc_init).all(): Loading Loading
src/borch/distributions/distribution_utils.py +13 −3 Original line number Diff line number Diff line Loading @@ -192,6 +192,12 @@ def _verify_is_continuous(rv): if not is_continuous(rv.support): raise RuntimeError(f"{rv} is not continuous") def _dist_mean(dist): try: mean = dist.mean.detach() except NotImplementedError: mean = dist.sample(sample_shape=torch.Size([1000])).mean(0) return mean def normal_distribution_from_rv( rv, log_scale: Union[Tensor, Number], loc_at_mean=True Loading @@ -217,7 +223,7 @@ def normal_distribution_from_rv( """ _verify_is_continuous(rv) if loc_at_mean: loc_init = rv.distribution.mean.detach() loc_init = _dist_mean(rv.distribution).detach() else: loc_init = rv.detach() if not torch.isfinite(loc_init).all(): Loading Loading @@ -263,7 +269,11 @@ def delta_distribution_from_rv(rv): def _scale(rv): if isinstance(rv, (dist.StudentT, dist.Cauchy)): return rv.scale return rv.distribution.stddev try: scale =rv.distribution.stddev except NotImplementedError: scale = rv.distribution.sample(sample_shape=torch.Size([1000])).std() return scale def scaled_normal_dist_from_rv(rv, scaling, loc_at_mean=True): Loading Loading @@ -291,7 +301,7 @@ def scaled_normal_dist_from_rv(rv, scaling, loc_at_mean=True): """ _verify_is_continuous(rv) if loc_at_mean: loc_init = rv.distribution.mean.detach() loc_init = _dist_mean(rv.distribution).detach() else: loc_init = rv.detach() if not torch.isfinite(loc_init).all(): Loading