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:
ModuleApply scaled dot-product multi-head attention.
Query, key, and value projections split
hidden_sizeintonum_headsheads. Attention is computed assoftmax(Q K^T / sqrt(head_dim)) Vand projected back tohidden_size. Passingcontextenables cross-attention; otherwise the module performs self-attention. Boolean masks use JAX semantics:Truepermits a query-key pair andFalseblocks it.axis_namesandpartition_specare optional mappings withq_proj,k_proj,v_proj, ando_projkeys.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)
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:
ModuleApply the Transformer’s position-wise feed-forward network.
Every sequence position is transformed independently by
Linear(hidden, intermediate), an activation, dropout, andLinear(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:
ModuleApply one Transformer encoder layer.
The original post-normalization order is used by default. Set
norm_first=Truefor 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:
ModuleApply a stack of Transformer encoder layers to batch-first inputs.
- 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:
ModuleApply one Transformer decoder layer.
The layer contains masked self-attention, encoder-decoder cross-attention, and a position-wise feed-forward network.
memory=Noneskips 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:
ModuleApply a stack of Transformer decoder layers to batch-first inputs.
- 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:
ModuleApply 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]. Setbatch_first=Truefor[batch, sequence, hidden].Attention masks are boolean arrays where
Truepermits attention. Key-padding masks use the inverse convention whereTruemarks an ignored key position. The target is causal by default, matching the autoregressive decoder in the original paper.With the original
norm_first=Falselayout, normalization occurs only after each residual sublayer and no extra stack-final norm is created. Thenorm_first=Truevariant adds a final norm after each complete encoder or decoder stack.axis_namesandpartition_specacceptq_proj,k_proj,v_proj,o_proj,input,output, andnormkeys.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)¶
- taktiny.nn.default_transformer_bias_initializer(key, shape, dtype=None, out_sharding=None)¶
An initializer that returns a constant array full of zeros.
The
keyargument is ignored.- Return type:
- 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)