Skip to main content

Vanilla Transformer

PyPI version

JAX/Flax implimentation of 'Attention Is All You Need' by Vaswani et al. (https://arxiv.org/abs/1706.03762)

Installation

Use the package manager pip to install the package in the following way:

pip install vanilla-transformer-jax

Usage

To use the entire Transformer model (encoder and decoder), you can use the following way:

from jax import random
from vtransformer import Transformer # imports Transformer class

model = Transformer() # model hyperparameters can be tuned, otherwise defualts mentioned in paper shall be used

prng = random.PRNGKey(42)

example_input_src = jax.random.randint(prng, (3,4), minval=0, maxval=10000)
example_input_trg = jax.random.randint(prng, (3,5), minval=0, maxval=10000)
mask = jax.array([1, 1, 1, 0, 0])

init = model.init(prng, example_input_src, example_input_trg, mask) #initializing the params of model

output = model.apply(init, example_input_src, example_input_trg, mask) # getting output

To use Encoder and Decoder seperately, you can do so in the following way:

encoding = model.encoder(init, example_input_src)  #using only the encoder
decoding = model.decoder(init, example_input_trg, encoding, mask) #using only the decoder

Contributing

This library is not perfect and can be improved in quite a few factors.

Pull requests are welcome. For major changes, please open an issue first to discuss what you would like to change.

Please make sure to update tests as appropriate.

License

MIT

Metadata

Release files for vanilla-transformer-jax 0.0.4

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for vanilla-transformer-jax 0.0.4
File Size Uploaded
vanilla-transformer-jax-0.0.4.tar.gz 4.4 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for vanilla-transformer-jax 0.0.4
File Interpreter ABI Platform
vanilla_transformer_jax-0.0.4-py3-none-any.whl Python 3 none any Details

Total release size: 9.4 kB

Release files / vanilla-transformer-jax-0.0.4.tar.gz

Download URL vanilla-transformer-jax-0.0.4.tar.gz
Size 4.4 kB
Tags Source
SHA-256 checksum
How to use checksums
f0e09c5d00f850507dc64dc402517f553f4339b2f2aba3a6cd4f21ca5e72778d
BLAKE2b-256 checksum
How to use checksums
94da354beea34817dea93dbf4c83b0b46f3d828501568f02a61d3ca80bf57dc9
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/3.7.1 importlib_metadata/3.7.3 pkginfo/1.8.2 requests/2.25.1 requests-toolbelt/0.9.1 tqdm/4.59.0 CPython/3.9.1

Release files / vanilla_transformer_jax-0.0.4-py3-none-any.whl

Download URL vanilla_transformer_jax-0.0.4-py3-none-any.whl
Size 5.0 kB
Tags Python 3
SHA-256 checksum
How to use checksums
8762fd684a98a58b69a4be3770b0e4bdd59f8e3dbcfe9161a9c4f870520beb8b
BLAKE2b-256 checksum
How to use checksums
ba3b763963abc597446f17128fa0a3348adb80d76c7236eae0025ccf265cba60
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/3.7.1 importlib_metadata/3.7.3 pkginfo/1.8.2 requests/2.25.1 requests-toolbelt/0.9.1 tqdm/4.59.0 CPython/3.9.1
Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page