google / google/flax

Adding transformer encoder and decoder layers to flax source as in pytorch

Open
#5,176 3 comments 0 reactions 0 assignees View on GitHub
Dominant language
Jupyter Notebook
Stars
7.3k
Forks
833
Avg merge
5h 11m
Merged PRs (30d)
5

Description

The pytorch source consists of implementations of wrappers for transformer modules.

SRC: https://github.com/pytorch/pytorch/blob/v2.9.1/torch/nn/modules/transformer.py#L966

I want to add such implementation for ease of use / ux. I will make a new file: `flax/flax/nnx/nn/transformer.py` which will contain the following modules:
* `TransformerEncoderLayer`
* `TransformerEncoder`
* `TransformerDecoderLayer`
* `TransformerDecoder`
* `Transformer`

I will keep it consistent with `nnx.Linear` and `nnx.MultiHeadAttention` modules and update the docs too, if needed, I can implement custom separate attentions such as `MHSA`(for full) and `GQA`(with kv-cache) based on review. Can I do a PR? @cgarciae @vfdev-5

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.