Skip to content

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

Description

@coder0143

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

Activity

  1. samanklesaria commented on Jan 7, 2026

    @samanklesaria
    Collaborator

    The comment for PyTorch's TransformerDecoder says the following:

    The intent of this layer is as a reference implementation for foundational understanding
    and thus it contains only limited features relative to newer Transformer architectures.
    Given the fast pace of innovation in transformer-like architectures, we recommend
    exploring this tutorial <https://pytorch.org/tutorials/intermediate/transformer_building_blocks.html>_
    to build efficient layers from building blocks in core or using higher
    level libraries from the PyTorch Ecosystem <https://landscape.pytorch.org/>_.

    That is to say, these layers in PyTorch are primarily for the sake of pedagogy. If the official recommendations from PyTorch are to "build efficient layers from building blocks in core" instead, it seems like that's the direction we should steer flax users as well.

    It seems like GQA is already possible in jax if the number of query heads is different from the number of key/value heads (see jax.nn.dot_product_attention). But flax.nnx.dot_product_attention doesn't inherit this functionality. I'll open a fix for that.

  2. coder0143 commented on Jan 8, 2026

    @coder0143
    Author

    Can I do a PR for it, can you review it for me @samanklesaria

  3. samanklesaria commented on Jan 8, 2026

    @samanklesaria
    Collaborator

    @coder0143 looks like @ayulockedin beat you to it with #5180

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions