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
25 changes: 23 additions & 2 deletions pyro/ops/provenance.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,13 @@
_Tensor = TypeVar("_Tensor", bound=torch.Tensor)


def _is_size(x: object) -> bool:
# ``torch.Size`` subclasses ``tuple``, so the pytree helpers below would
# otherwise flatten it and rebuild it as a plain tuple. Treating it as a
# leaf keeps the type intact.
return isinstance(x, torch.Size)


class ProvenanceTensor(torch.Tensor):
"""
Provenance tracking implementation in Pytorch.
Expand Down Expand Up @@ -85,6 +92,13 @@ def track_provenance(x, provenance: frozenset):
track_provenance.register(torch.Tensor)(ProvenanceTensor)


@track_provenance.register(torch.Size)
def _track_provenance_size(x, provenance: frozenset):
# A torch.Size holds plain ints, so there is no provenance to add; return
# it unchanged rather than letting the tuple handler rebuild it.
return x


@track_provenance.register(frozenset)
@track_provenance.register(set)
def _track_provenance_set(x, provenance: frozenset):
Expand All @@ -95,7 +109,9 @@ def _track_provenance_set(x, provenance: frozenset):
@track_provenance.register(tuple)
@track_provenance.register(dict)
def _track_provenance_pytree(x, provenance: frozenset):
return tree_map(partial(track_provenance, provenance=provenance), x)
return tree_map(
partial(track_provenance, provenance=provenance), x, is_leaf=_is_size
)


@track_provenance.register
Expand Down Expand Up @@ -123,6 +139,11 @@ def _extract_provenance_tensor(x):
return x._t, x._provenance


@extract_provenance.register(torch.Size)
def _extract_provenance_size(x):
return x, frozenset()


@extract_provenance.register(frozenset)
@extract_provenance.register(set)
def _extract_provenance_set(x):
Expand All @@ -140,7 +161,7 @@ def _extract_provenance_set(x):
@extract_provenance.register(tuple)
@extract_provenance.register(dict)
def _extract_provenance_pytree(x):
flat_args, spec = tree_flatten(x)
flat_args, spec = tree_flatten(x, is_leaf=_is_size)
xs = []
provenance = frozenset()
for x, p in map(extract_provenance, flat_args):
Expand Down
61 changes: 60 additions & 1 deletion tests/ops/test_provenance.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,16 @@
import pytest
import torch

from pyro.ops.provenance import ProvenanceTensor, get_provenance, track_provenance
import pyro
import pyro.distributions as dist
from pyro.infer.inspect import get_model_relations
from pyro.ops.provenance import (
ProvenanceTensor,
detach_provenance,
extract_provenance,
get_provenance,
track_provenance,
)
from tests.common import assert_equal, requires_cuda


Expand Down Expand Up @@ -66,3 +75,53 @@ def test_track_provenance(x):
old_provenance = get_provenance(x)
provenance = old_provenance | new_provenance
assert provenance == get_provenance(track_provenance(x, new_provenance))


@pytest.mark.parametrize(
"x",
[
torch.Size([3]),
torch.Size([2, 5]),
torch.Size([]),
[(torch.Size([3]),), {}],
[(), {"shape": torch.Size([2, 5])}],
[(torch.zeros(2), torch.Size([3])), {}],
],
ids=["size", "size_2d", "size_empty", "in_args", "in_kwargs", "mixed_with_tensor"],
)
def test_provenance_preserves_torch_size(x):
"""torch.Size subclasses tuple, so it must not be rebuilt as a plain tuple."""

def assert_sizes_intact(original, result):
if isinstance(original, torch.Tensor):
return # tensors are legitimately wrapped/unwrapped
assert type(original) is type(result)
if isinstance(original, torch.Size):
assert result == original
elif isinstance(original, (list, tuple)):
for a, b in zip(original, result):
assert_sizes_intact(a, b)
elif isinstance(original, dict):
for key in original:
assert_sizes_intact(original[key], result[key])

assert_sizes_intact(x, track_provenance(x, frozenset({"a"})))
assert_sizes_intact(x, extract_provenance(x)[0])
assert_sizes_intact(x, detach_provenance(x))


def test_multinomial_render_model():
"""https://github.com/pyro-ppl/pyro/issues/3436"""

def model():
probs = pyro.param("probs", torch.tensor([0.1, 0.2, 0.7]))
return pyro.sample("y", dist.Multinomial(total_count=10, probs=probs))

relations = get_model_relations(model)
assert "y" in relations["sample_sample"] or "y" in relations["sample_param"]


def test_provenance_tensor_new_with_size():
"""Tensor.new() reads a torch.Size as a shape but a tuple as data."""
x = ProvenanceTensor(torch.zeros(4, dtype=torch.long), frozenset({"x"}))
assert tuple(x.new(torch.Size([3])).shape) == (3,)