Transfusion in Pytorch
Project description
Transfusion - Pytorch (wip)
Pytorch implementation of Transfusion, "Predict the Next Token and Diffuse Images with One Multi-Modal Model", from MetaAI.
Once completed, will also extend this to flow matching, as well as audio, video, perhaps even policies.
Install
$ pip install transfusion-pytorch
Usage
One modality, say images
from torch import randint, randn
from transfusion_pytorch import Transfusion
model = Transfusion(
num_text_tokens = 256,
dim_latent = 384,
transformer = dict(
dim = 512,
depth = 8
)
)
text_and_images = [
[randint(0, 256, (16,)), randn(4, 384), randint(0, 256, (8,)), randn(6, 384)],
[randint(0, 256, (16,)), randn(7, 384), randint(0, 256, (5,)), randn(2, 384), randint(0, 256, (9,))]
]
loss = model(text_and_images)
loss.backward()
Multiple different modalities
from torch import randint, randn
from transfusion_pytorch import Transfusion
model = Transfusion(
num_text_tokens = 256,
dim_latent = (384, 192), # specify multiple latent dimensions
transformer = dict(
dim = 512,
depth = 8
)
)
# then for the Tensors of type float, you can pass a tuple[int, Tensor] and specify the modality index in the first position
text_images_and_audio = [
[randint(0, 256, (16,)), (0, randn(4, 384)), randint(0, 256, (8,)), (1, randn(6, 192))],
[randint(0, 256, (16,)), randn(7, 384), randint(0, 256, (5,)), (1, randn(2, 192)), randint(0, 256, (9,))]
]
loss = model(text_images_and_audio)
loss.backward()
Citations
@inproceedings{Zhou2024TransfusionPT,
title = {Transfusion: Predict the Next Token and Diffuse Images with One Multi-Modal Model},
author = {Chunting Zhou and Lili Yu and Arun Babu and Kushal Tirumala and Michihiro Yasunaga and Leonid Shamis and Jacob Kahn and Xuezhe Ma and Luke Zettlemoyer and Omer Levy},
year = {2024},
url = {https://api.semanticscholar.org/CorpusID:271909855}
}
@misc{Rubin2024,
author = {Ohad Rubin},
url = {https://medium.com/@ohadrubin/exploring-weight-decay-in-layer-normalization-challenges-and-a-reparameterization-solution-ad4d12c24950}
}
Project details
Release history Release notifications | RSS feed
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
transfusion_pytorch-0.0.16.tar.gz
(347.5 kB
view hashes)
Built Distribution
Close
Hashes for transfusion_pytorch-0.0.16.tar.gz
Algorithm | Hash digest | |
---|---|---|
SHA256 | 6b75f16389005c2209c8583e8148235ef60c46fba3f1b08e8dc154530e3bd31d |
|
MD5 | b49cdbbb40f2b3fbef6c01467a39ee2c |
|
BLAKE2b-256 | 44d7eb03b51caf6492748fd4ab98cd1eb628e5dd7ae7489a902e54ccc017dd2b |
Close
Hashes for transfusion_pytorch-0.0.16-py3-none-any.whl
Algorithm | Hash digest | |
---|---|---|
SHA256 | aea15d4eda39b0995114491f36e1d7f5b1af9a0efa63aa285101167c158f7317 |
|
MD5 | 1cb4a20f55fb17eaa14fa3cd6eee77ea |
|
BLAKE2b-256 | cfbcfdbbd7dfc5f2aceb5aab29f82e733403a5ce6ca8e1aa42c7657669c42a01 |