Transformer

Attention mechanisms, feed-forward networks, and transformer stacks.

class taktiny.nn.Attention(hidden_size, num_heads, head_dim=None, *, context_dim=None, dropout=0.0, scaling=None, bias=True, dtype=None, rngs, kernel_initializer=<function variance_scaling.<locals>.init>, bias_initializer=<function zeros>, quant=None, dot_general=None, axis_names=None, partition_spec=None, kernel_metadata=None, bias_metadata=None, precision=None, preferred_element_type=None)[source]

Bases: Module

Apply scaled dot-product multi-head attention.

Query, key, and value projections split hidden_size into num_heads heads. Attention is computed as softmax(Q K^T / sqrt(head_dim)) V and projected back to hidden_size. Passing context enables cross-attention; otherwise the module performs self-attention. Boolean masks use JAX semantics: True permits a query-key pair and False blocks it.

axis_names and partition_spec are optional mappings with q_proj, k_proj, v_proj, and o_proj keys.

Example

>>> import jax.numpy as jnp
>>> from taktiny import nn
>>> attention = nn.Attention(8, 2, rngs=nn.Rngs(0), dropout=0.0)
>>> attention(jnp.ones((3, 5, 8))).shape
(3, 5, 8)
Reference:

Ashish Vaswani et al., “Attention Is All You Need” (2017), https://arxiv.org/abs/1706.03762

Parameters:
  • hidden_size (int)

  • num_heads (int)

  • head_dim (int | None)

  • context_dim (int | None)

  • dropout (float)

  • scaling (float | None)

  • bias (bool | Sequence[bool])

  • dtype (DType | None)

  • rngs (Rngs)

  • kernel_initializer (Initializer)

  • bias_initializer (Initializer)

  • quant (QuantConfig)

  • dot_general (DotGeneral | None)

  • axis_names (AxisNamesMap | None)

  • partition_spec (PartitionSpecMap | None)

  • kernel_metadata (MetaData | None)

  • bias_metadata (MetaData | None)

  • precision (PrecisionLike)

  • preferred_element_type (DTypeLike | None)

class taktiny.nn.FeedForward(hidden_size, intermediate_size, *, activation='relu', dropout=0.0, bias=True, dtype=None, rngs, kernel_initializer=<function variance_scaling.<locals>.init>, bias_initializer=<function zeros>, quant=None, dot_general=None, axis_names=None, partition_spec=None, kernel_metadata=None, bias_metadata=None, precision=None, preferred_element_type=None)[source]

Bases: Module

Apply the Transformer’s position-wise feed-forward network.

Every sequence position is transformed independently by Linear(hidden, intermediate), an activation, dropout, and Linear(intermediate, hidden).

Example

>>> import jax.numpy as jnp
>>> from taktiny import nn
>>> feed_forward = nn.FeedForward(8, 32, rngs=nn.Rngs(0), dropout=0.0)
>>> feed_forward(jnp.ones((2, 5, 8))).shape
(2, 5, 8)
Parameters:
  • hidden_size (int)

  • intermediate_size (int)

  • activation (Activation)

  • dropout (float)

  • bias (bool)

  • dtype (DType | None)

  • rngs (Rngs)

  • kernel_initializer (Initializer)

  • bias_initializer (Initializer)

  • quant (QuantConfig)

  • dot_general (DotGeneral | None)

  • axis_names (AxisNamesMap | None)

  • partition_spec (PartitionSpecMap | None)

  • kernel_metadata (MetaData | None)

  • bias_metadata (MetaData | None)

  • precision (PrecisionLike)

  • preferred_element_type (DTypeLike | None)

class taktiny.nn.TransformerEncoderLayer(hidden_size, num_heads, intermediate_size, *, dropout=0.1, activation='relu', norm_first=False, norm_eps=1e-05, bias=True, dtype=None, rngs, kernel_initializer=<function variance_scaling.<locals>.init>, bias_initializer=<function zeros>, quant=None, dot_general=None, axis_names=None, partition_spec=None, kernel_metadata=None, bias_metadata=None, precision=None, preferred_element_type=None)[source]

Bases: Module

Apply one Transformer encoder layer.

The original post-normalization order is used by default. Set norm_first=True for the common pre-normalization variant.

Example

>>> import jax.numpy as jnp
>>> from taktiny import nn
>>> layer = nn.TransformerEncoderLayer(8, 2, 32, rngs=nn.Rngs(0), dropout=0.0)
>>> layer(jnp.ones((2, 5, 8))).shape
(2, 5, 8)
Parameters:
  • hidden_size (int)

  • num_heads (int)

  • intermediate_size (int)

  • dropout (float)

  • activation (Activation)

  • norm_first (bool)

  • norm_eps (float)

  • bias (bool)

  • dtype (DType | None)

  • rngs (Rngs)

  • kernel_initializer (Initializer)

  • bias_initializer (Initializer)

  • quant (QuantConfig)

  • dot_general (DotGeneral | None)

  • axis_names (AxisNamesMap | None)

  • partition_spec (PartitionSpecMap | None)

  • kernel_metadata (MetaData | None)

  • bias_metadata (MetaData | None)

  • precision (PrecisionLike)

  • preferred_element_type (DTypeLike | None)

class taktiny.nn.TransformerEncoder(layers, norm=None)[source]

Bases: Module

Apply a stack of Transformer encoder layers to batch-first inputs.

Parameters:
class taktiny.nn.TransformerDecoderLayer(hidden_size, num_heads, intermediate_size, *, dropout=0.1, activation='relu', norm_first=False, norm_eps=1e-05, bias=True, dtype=None, rngs, kernel_initializer=<function variance_scaling.<locals>.init>, bias_initializer=<function zeros>, quant=None, dot_general=None, axis_names=None, partition_spec=None, kernel_metadata=None, bias_metadata=None, precision=None, preferred_element_type=None)[source]

Bases: Module

Apply one Transformer decoder layer.

The layer contains masked self-attention, encoder-decoder cross-attention, and a position-wise feed-forward network. memory=None skips the cross-attention branch, allowing decoder-only use.

Example

>>> import jax.numpy as jnp
>>> from taktiny import nn
>>> layer = nn.TransformerDecoderLayer(8, 2, 32, rngs=nn.Rngs(0), dropout=0.0)
>>> layer(jnp.ones((2, 4, 8)), jnp.ones((2, 6, 8)), self_is_causal=True).shape
(2, 4, 8)
Parameters:
  • hidden_size (int)

  • num_heads (int)

  • intermediate_size (int)

  • dropout (float)

  • activation (Activation)

  • norm_first (bool)

  • norm_eps (float)

  • bias (bool)

  • dtype (DType | None)

  • rngs (Rngs)

  • kernel_initializer (Initializer)

  • bias_initializer (Initializer)

  • quant (QuantConfig)

  • dot_general (DotGeneral | None)

  • axis_names (AxisNamesMap | None)

  • partition_spec (PartitionSpecMap | None)

  • kernel_metadata (MetaData | None)

  • bias_metadata (MetaData | None)

  • precision (PrecisionLike)

  • preferred_element_type (DTypeLike | None)

class taktiny.nn.TransformerDecoder(layers, norm=None)[source]

Bases: Module

Apply a stack of Transformer decoder layers to batch-first inputs.

Parameters:
class taktiny.nn.Transformer(hidden_size=512, num_heads=8, num_encoder_layers=6, num_decoder_layers=6, intermediate_size=2048, dropout=0.1, activation='relu', *, norm_eps=1e-05, batch_first=False, norm_first=False, bias=True, custom_encoder=None, custom_decoder=None, dtype=None, rngs, kernel_initializer=<function variance_scaling.<locals>.init>, bias_initializer=<function zeros>, quant=None, dot_general=None, axis_names=None, partition_spec=None, kernel_metadata=None, bias_metadata=None, precision=None, preferred_element_type=None)[source]

Bases: Module

Apply the encoder-decoder architecture from Attention Is All You Need.

Inputs are embedded hidden states; token embeddings and positional encodings remain the caller’s responsibility. By default, unbatched inputs use [sequence, hidden] and batched inputs use [sequence, batch, hidden]. Set batch_first=True for [batch, sequence, hidden].

Attention masks are boolean arrays where True permits attention. Key-padding masks use the inverse convention where True marks an ignored key position. The target is causal by default, matching the autoregressive decoder in the original paper.

With the original norm_first=False layout, normalization occurs only after each residual sublayer and no extra stack-final norm is created. The norm_first=True variant adds a final norm after each complete encoder or decoder stack.

axis_names and partition_spec accept q_proj, k_proj, v_proj, o_proj, input, output, and norm keys.

Example

>>> import jax.numpy as jnp
>>> from taktiny import nn
>>> model = nn.Transformer(
...     hidden_size=8,
...     num_heads=2,
...     num_encoder_layers=1,
...     num_decoder_layers=1,
...     intermediate_size=32,
...     dropout=0.0,
...     batch_first=True,
...     rngs=nn.Rngs(0),
... )
>>> model(jnp.ones((2, 5, 8)), jnp.ones((2, 3, 8))).shape
(2, 3, 8)
Reference:

Ashish Vaswani et al., “Attention Is All You Need” (2017), https://arxiv.org/abs/1706.03762

Parameters:
  • hidden_size (int)

  • num_heads (int)

  • num_encoder_layers (int)

  • num_decoder_layers (int)

  • intermediate_size (int)

  • dropout (float)

  • activation (Activation)

  • norm_eps (float)

  • batch_first (bool)

  • norm_first (bool)

  • bias (bool)

  • custom_encoder (Module | None)

  • custom_decoder (Module | None)

  • dtype (DType | None)

  • rngs (Rngs)

  • kernel_initializer (Initializer)

  • bias_initializer (Initializer)

  • quant (QuantConfig)

  • dot_general (DotGeneral | None)

  • axis_names (AxisNamesMap | None)

  • partition_spec (PartitionSpecMap | None)

  • kernel_metadata (MetaData | None)

  • bias_metadata (MetaData | None)

  • precision (PrecisionLike)

  • preferred_element_type (DTypeLike | None)

taktiny.nn.default_transformer_initializer(key, shape, dtype=None, out_sharding=None)
Return type:

Array

Parameters:
taktiny.nn.default_transformer_bias_initializer(key, shape, dtype=None, out_sharding=None)

An initializer that returns a constant array full of zeros.

The key argument is ignored.

Return type:

Array

Parameters:
>>> import jax, jax.numpy as jnp
>>> jax.nn.initializers.zeros(jax.random.key(42), (2, 3), jnp.float32)
Array([[0., 0., 0.],
       [0., 0., 0.]], dtype=float32)