patrick-kidger / patrick-kidger/diffrax
can I make a continuous normalizing flow faster than real nvp?
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 2.1k
- Forks
- 189
- Avg merge
- 3d 18h
- Merged PRs (30d)
- 1
Description
I implemented the CNF example and added it as a head to a transformer to do inference over some continuous variables in a probabilistic program, however, it's wayyyyy slower than an my equinox implementation of real nvp inspired by this excelent tutorial.
However, the elegance of a CNF is just too good to just ignore, any recommendations?
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start with the linked continuous normalising flow example in the Diffrax documentation and compare it with src/real_nvp.py from the linked Equinox implementation. Measure the two approaches under the same inference setup, then determine whether a concrete Diffrax change can improve the CNF performance; done requires an agreed optimization or recommendation.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python
- Domain
- machine-learning, performance
- Issue type
- Feature
- Difficulty
- 5/5
- Estimated time
- Over a week
- Activity status
- Stale
- Clarity
- Needs clarification
- Newbie friendliness
- 20/100