facebookresearch / facebookresearch/flow_matching

Invalid t shape for data with ndim > 1 in GeodesicProbPath.sample()

Open
#73 0 comments 1 reaction 0 assignees View on GitHub
Dominant language
Python
Stars
4.7k
Forks
373
PR merge metrics
No merged PRs in 30d

Description

Hi and thanks a lot for your amazing work on this library.

**Describe the bug**
After upgrading from version `1.0.9` to `1.0.10`, a `RuntimeError` occurs when calling `GeodesicProbPath.sample()` with high-dimensional input (i.e., inputs with shape > 2D including batch dimension).
```
RuntimeError: einsum(): the number of subscripts in the equation
(1) does not match the number of dimensions (2) for operand 0 and no ellipsis was given
```

**To Reproduce**
```python
from flow_matching.path import GeodesicProbPath, PathSample
from flow_matching.path.scheduler import CondOTScheduler
from flow_matching.utils.manifolds import Euclidean
import torch

batch_size = 128
data_dim = (16, 4)

x_0 = torch.randn((batch_size, *data_dim)) # (128, 16, 4)
x_1 = torch.randn((batch_size, *data_dim)) # (128, 16, 4)
t = torch.linspace(0, 1, batch_size) # (128)

manifold = Euclidean()
scheduler = CondOTScheduler()
path = GeodesicProbPath(scheduler, manifold)

sample: PathSample = path.sample(x_0=x_0, x_1=x_1, t=t)
# RuntimeError: einsum(): the number of subscripts in the equation
# (1) does not match the number of dimensions (2) for operand 0 and no ellipsis was given
```

**Expected behavior**
The function should support data with arbitrary trailing dimensions (e.g., (batch_size, D1, D2)), not just 2D inputs.

Thank you in advance
Lukas

Contributor guide

Open the contributing guide

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.