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.17.tar.gz
(347.8 kB
view hashes)
Built Distribution
Close
Hashes for transfusion_pytorch-0.0.17.tar.gz
Algorithm | Hash digest | |
---|---|---|
SHA256 | 6b7fc0b7386db0bf61479beb28fefcca7c5d89c2b87f6e6c8601bed3af4d3d71 |
|
MD5 | bbff26347cde0e08796efa2ba4aeefd9 |
|
BLAKE2b-256 | 15c1b97a088362250e46fa44aa94c36c604cc56ed973701d0fc89470a7d584ee |
Close
Hashes for transfusion_pytorch-0.0.17-py3-none-any.whl
Algorithm | Hash digest | |
---|---|---|
SHA256 | 67580d4c0ac580c5df2f2a50186e030843c0d75a15820566ccc81dac75a4ad66 |
|
MD5 | dcb9923118f9fe5856424824ff218752 |
|
BLAKE2b-256 | 6be993c803d48ce316b3b9d73aa85b568769b1dd25bf7b1e8c1cff98e1bf75bd |