Skip to content
Open
Show file tree
Hide file tree
Changes from all 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
38 changes: 30 additions & 8 deletions pyro/distributions/transforms/spline_coupling.py
Original file line number Diff line number Diff line change
Expand Up @@ -166,7 +166,12 @@ def log_abs_det_jacobian(self, x, y):


def spline_coupling(
input_dim, split_dim=None, hidden_dims=None, count_bins=8, bound=3.0
input_dim,
split_dim=None,
hidden_dims=None,
count_bins=8,
bound=3.0,
order="linear",
):
"""
A helper function to create a
Expand All @@ -175,6 +180,18 @@ def spline_coupling(

:param input_dim: Dimension of input variable
:type input_dim: int
:param split_dim: Zero-indexed dimension :math:`d` upon which to perform input/
output split for transformation.
:type split_dim: int
:param hidden_dims: The dimensions of the hidden units of the hypernet.
:type hidden_dims: list
:param count_bins: The number of segments comprising the spline.
:type count_bins: int
:param bound: The quantity :math:`K` determining the bounding box,
:math:`[-K,K]\times[-K,K]`, of the spline.
:type bound: float
:param order: One of ['linear', 'quadratic'] specifying the order of the spline.
:type order: string

"""

Expand All @@ -184,15 +201,20 @@ def spline_coupling(
if hidden_dims is None:
hidden_dims = [input_dim * 10, input_dim * 10]

# Rational linear splines have an additional lambda parameter per bin,
# while quadratic splines only use widths, heights and derivatives.
param_dims = [
(input_dim - split_dim) * count_bins,
(input_dim - split_dim) * count_bins,
(input_dim - split_dim) * (count_bins - 1),
]
if order == "linear":
param_dims.append((input_dim - split_dim) * count_bins)

nn = DenseNN(
split_dim,
hidden_dims,
param_dims=[
(input_dim - split_dim) * count_bins,
(input_dim - split_dim) * count_bins,
(input_dim - split_dim) * (count_bins - 1),
(input_dim - split_dim) * count_bins,
],
param_dims=param_dims,
)

return SplineCoupling(input_dim, split_dim, nn, count_bins, bound)
return SplineCoupling(input_dim, split_dim, nn, count_bins, bound, order)
39 changes: 36 additions & 3 deletions pyro/poutine/replay_messenger.py
Original file line number Diff line number Diff line change
@@ -1,17 +1,28 @@
# Copyright (c) 2017-2019 Uber Technologies, Inc.
# SPDX-License-Identifier: Apache-2.0

from typing import TYPE_CHECKING, Dict, Optional
import warnings
from typing import TYPE_CHECKING, Any, Dict, Optional

import torch

from pyro.poutine.messenger import Messenger
from pyro.poutine.util import site_is_subsample

if TYPE_CHECKING:
import torch

from pyro.poutine.runtime import Message
from pyro.poutine.trace_struct import Trace


def _subsample_values_equal(a: Any, b: Any) -> bool:
"""Compare two subsample index values for equality."""
if isinstance(a, torch.Tensor) and isinstance(b, torch.Tensor):
return torch.equal(a, b)
if isinstance(a, torch.Tensor) or isinstance(b, torch.Tensor):
return False
return bool(a == b)


class ReplayMessenger(Messenger):
"""
Given a callable that contains Pyro primitive calls,
Expand Down Expand Up @@ -76,6 +87,28 @@ def _pyro_sample(self, msg: "Message") -> None:
return None
if guide_msg["type"] != "sample" or guide_msg["is_observed"]:
raise RuntimeError("site {} must be sampled in trace".format(name))
# Warn when replaying a subsample site whose explicit value differs
# from the guide's independently drawn subsample. This happens when a
# model passes an explicit ``subsample=idx`` to ``pyro.plate`` but is
# composed with a guide that draws its own subsample (e.g. an
# ``AutoGuide`` built without ``create_plates=``). Silently overriding
# the model's index decouples the model's and guide's minibatches.
# See https://github.com/pyro-ppl/pyro/issues/3468
if (
msg["value"] is not None
and guide_msg["value"] is not None
and site_is_subsample(msg)
and not _subsample_values_equal(msg["value"], guide_msg["value"])
):
warnings.warn(
"Replaying the subsample site '{}' with a value that differs "
"from the model's explicit subsample. If the model passes an "
"explicit ``subsample=idx`` to ``pyro.plate``, use a guide that "
"reuses the same subsample (e.g. pass ``create_plates=`` to "
"``AutoGuide``) so the model's and guide's minibatches stay "
"aligned.".format(name),
stacklevel=2,
)
msg["done"] = True
msg["value"] = guide_msg["value"]
msg["infer"] = guide_msg["infer"]
Expand Down
3 changes: 2 additions & 1 deletion tests/distributions/test_transforms.py
Original file line number Diff line number Diff line change
Expand Up @@ -419,7 +419,8 @@ def test_spline(self):
self._test(partial(T.spline, order=order))

def test_spline_coupling(self):
self._test(T.spline_coupling)
for order in ["linear", "quadratic"]:
self._test(partial(T.spline_coupling, order=order))

def test_spline_autoregressive(self):
self._test(T.spline_autoregressive)
Expand Down
50 changes: 50 additions & 0 deletions tests/poutine/test_poutines.py
Original file line number Diff line number Diff line change
Expand Up @@ -123,6 +123,56 @@ def test_replay_full_repeat(self):
assert_equal(model_trace.nodes[name]["value"], tr2.nodes[name]["value"])


def test_replay_warns_on_mismatched_explicit_subsample():
"""Replay should warn when it overrides a model's explicit subsample.

When a model passes an explicit ``subsample=idx`` to ``pyro.plate`` but is
composed with a guide that draws its own subsample (e.g. an ``AutoGuide``
without ``create_plates=``), replay silently overrode the model's index with
the guide's. We now warn about this decoupling (see #3468).
"""
from pyro.infer.autoguide import AutoNormal

N = 6
idx = torch.tensor([4, 0, 1])

def model(idx):
with pyro.plate("data", N, dim=-1, subsample=idx):
pyro.sample("mu", dist.Normal(0.0, 1.0))

pyro.clear_param_store()
guide = AutoNormal(model)
guide(idx)
guide_trace = poutine.trace(guide).get_trace(idx)

with pytest.warns(UserWarning, match="subsample site 'data'"):
poutine.trace(poutine.replay(model, trace=guide_trace)).get_trace(idx)


def test_replay_no_warn_when_subsample_matches():
"""Replay should not warn when the guide reuses the model's subsample."""
from pyro.infer.autoguide import AutoNormal

N = 6
idx = torch.tensor([4, 0, 1])

def create_plates(idx):
return pyro.plate("data", N, dim=-1, subsample=idx)

def model(idx):
with pyro.plate("data", N, dim=-1, subsample=idx):
pyro.sample("mu", dist.Normal(0.0, 1.0))

pyro.clear_param_store()
guide = AutoNormal(model, create_plates=create_plates)
guide(idx)
guide_trace = poutine.trace(guide).get_trace(idx)

with warnings.catch_warnings():
warnings.simplefilter("error")
poutine.trace(poutine.replay(model, trace=guide_trace)).get_trace(idx)


class BlockHandlerTests(NormalNormalNormalHandlerTestCase):
def test_block_hide_fn(self):
model_trace = poutine.trace(
Expand Down