Reinterpret event_shape as batch_shape in torch.distributions

Viewed 101

Let's say I have a Gamma distribution

q = D.Independent(D.Gamma(alpha, beta), 1)

Where alpha and beta are the parameters vectors with dimensions [n_samples, batch_size, d].

The distribution will have for shapes:

print(q.batch_shape, q.event_shape)
>>> [n_samples, batch_size], [d]

the D.Independent wrapper class allows to reinterpret the batch dimensions as event dimensions (as done when initialising the Gamma distribution), but is there a way to do the opposite, that is:

Is there a way to reinterpret event dimensions as batch dimensions?

I know that I can re-define a new distribution, i.e

alpha = q.base_dist.concentration
beta = q.base_dist.rate
q = D.Gamma(alpha, beta)

that will have the correct dimensions

print(q.batch_shape, q.event_shape)
>>> [n_samples, batch_size, d], []

But that seems highly unpractical and inefficient.

Any ideas?

0 Answers
Related