pyro-ppl / pyro-ppl/numpyro

Vectorized interpretation for pyro.scan

Open
#686 3 comments 2 reactions 1 assignee View on GitHub

Nobody has claimed this yet.

enhancement
Dominant language
Python
Stars
2.8k
Forks
315
Avg merge
3d 9h
Merged PRs (30d)
27

Description

This proposes to implement a vectorized interpretation of pyro.scan that completely parallelizes over the time axis. This follows a hand-implementation of vectorization in pyro.contrib.epidemiology (in CompartmentalModel._relaxed_model() and ._vectorized_model. This interpretation works only under replay or condition (the standard interpretation of sampling cannot be time-parallelized), therefore this will require three-way interaction between pyro.scan, poutine.condition, and the new interpretation.

Here is a vague sketch of the parallel implementation of pyro.scan:

def vectorized_scan(transition, time, init):
    """
    This assumes ``init`` is a dict mapping unindexed sample site name (i.e.
    "x" rather than "x_0") to value. TODO generalize to PyTree.
    This assumes ``time`` is a range object. TODO generalize to jnp.arange?
    """
    # Trace the first step and assume model structure is fixed.
    # In pyro.contrib.epideiology we do this once at the start of inference.
    # Maybe we could memoize to avoid duplicated execution?
    with poutine.block(), poutine.trace() as tr:
        t = 0
        transition(init, t)
    names = [name for name in tr.trace.stochastic_nodes
             if name.endswith("_0")]

    # The remainder is vectorized over time.
    with pyro.plate("time", len(time)):  # or maybe jax.vmap
        t = slice(0, len(time), 1)  # or maybe jnp.arange

        # Record vectorized values.
        curr = {}
        prev = {}
        with poutine.block_trace_but_allow_replay_and_condition():
            for name in names:
                name_0 = "{}_{}".format(name, 0)
                name_t = "{}_{}".format(name, t)
                site_0 = tr.nodes[name_0]
                curr[name] = pyro.sample(name_t, site_0["fn"])
                prev[name] = torch.cat([site_0["value"].unsqueeze(0), curr[name]])

        # Execute vectorized transition.
        transition(prev, t)

    return ...

cc @fehiepsi @eb8680

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.