Commit 0566fb96 authored by Johan Gudmundsson's avatar Johan Gudmundsson
Browse files

make the dist from rv handle missing mean and std

parent d2328b5c
Loading
Loading
Loading
Loading
+13 −3
Original line number Diff line number Diff line
@@ -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
@@ -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():
@@ -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):
@@ -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():