encoder-decoder-transformer 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,27 @@
1
+ """Encoder-decoder transformer implemented from scratch in PyTorch."""
2
+
3
+ from .decoder import Decoder, DecoderLayer
4
+ from .embeddings import Embeddings
5
+ from .encoder import Encoder, EncoderLayer
6
+ from .encoder_decoder import EncoderDecoder
7
+ from .feed_forward import FeedForward
8
+ from .make_model import make_model
9
+ from .multi_head_attention import MultiHeadedAttention, attention
10
+ from .utils import Generator, LayerNorm, SublayerConnection, clones
11
+
12
+ __all__ = [
13
+ "Decoder",
14
+ "DecoderLayer",
15
+ "Embeddings",
16
+ "Encoder",
17
+ "EncoderDecoder",
18
+ "EncoderLayer",
19
+ "FeedForward",
20
+ "Generator",
21
+ "LayerNorm",
22
+ "MultiHeadedAttention",
23
+ "SublayerConnection",
24
+ "attention",
25
+ "clones",
26
+ "make_model",
27
+ ]
@@ -0,0 +1,129 @@
1
+ """Transformer decoder layers and stack."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import TYPE_CHECKING, cast
6
+
7
+ from torch import Tensor, nn
8
+
9
+ from .utils import LayerNorm, SublayerConnection, clones
10
+
11
+ if TYPE_CHECKING:
12
+ from .feed_forward import FeedForward
13
+ from .multi_head_attention import MultiHeadedAttention
14
+
15
+
16
+ class DecoderLayer(nn.Module):
17
+ """A single decoder layer with self-attention, cross-attention and FFN.
18
+
19
+ Args:
20
+ d_model: Feature dimension of the input.
21
+ self_attn: Multi-head self-attention module.
22
+ cross_attn: Multi-head cross-attention module.
23
+ ff: Position-wise feed-forward network.
24
+ dropout: Dropout probability used inside the sublayer connections.
25
+ """
26
+
27
+ def __init__(
28
+ self,
29
+ d_model: int,
30
+ self_attn: MultiHeadedAttention,
31
+ cross_attn: MultiHeadedAttention,
32
+ ff: FeedForward,
33
+ dropout: float = 0.1,
34
+ ) -> None:
35
+ """Initialize the attention and feed-forward sublayers.
36
+
37
+ Args:
38
+ d_model: Feature dimension of the input.
39
+ self_attn: Multi-head self-attention module.
40
+ cross_attn: Multi-head cross-attention module.
41
+ ff: Position-wise feed-forward network.
42
+ dropout: Dropout probability used inside the sublayer connections.
43
+ """
44
+ super().__init__()
45
+ self.self_attn: MultiHeadedAttention = self_attn
46
+ self.cross_attn: MultiHeadedAttention = cross_attn
47
+ self.ff: FeedForward = ff
48
+ self.sublayers: nn.ModuleList = nn.ModuleList(
49
+ [
50
+ SublayerConnection(d_model, dropout),
51
+ SublayerConnection(d_model, dropout),
52
+ SublayerConnection(d_model, dropout),
53
+ ]
54
+ )
55
+
56
+ def forward(
57
+ self,
58
+ x: Tensor,
59
+ encoder_output: Tensor,
60
+ src_mask: Tensor,
61
+ tgt_mask: Tensor,
62
+ ) -> Tensor:
63
+ """Apply self-attention, cross-attention and the feed-forward network.
64
+
65
+ Args:
66
+ x: Target input tensor of shape ``(batch, tgt_seq, d_model)``.
67
+ encoder_output: Encoder output of shape
68
+ ``(batch, src_seq, d_model)``.
69
+ src_mask: Source attention mask broadcastable to
70
+ ``(batch, 1, tgt_seq, src_seq)``.
71
+ tgt_mask: Target attention mask broadcastable to
72
+ ``(batch, 1, tgt_seq, tgt_seq)``.
73
+
74
+ Returns:
75
+ A tensor of shape ``(batch, tgt_seq, d_model)``.
76
+ """
77
+ x = self.sublayers[0](x, lambda t: self.self_attn(t, t, t, tgt_mask))
78
+ x = self.sublayers[1](
79
+ x, lambda t: self.cross_attn(t, encoder_output, encoder_output, src_mask)
80
+ )
81
+ x = self.sublayers[2](x, self.ff)
82
+ return x
83
+
84
+
85
+ class Decoder(nn.Module):
86
+ """Stack of ``N`` decoder layers followed by layer normalization.
87
+
88
+ Args:
89
+ decoder_layer: The decoder layer to clone.
90
+ n: Number of decoder layers in the stack.
91
+ d_model: Feature dimension of the input.
92
+ """
93
+
94
+ def __init__(self, decoder_layer: DecoderLayer, n: int, d_model: int) -> None:
95
+ """Initialize the cloned layers and final layer norm.
96
+
97
+ Args:
98
+ decoder_layer: The decoder layer to clone.
99
+ n: Number of decoder layers in the stack.
100
+ d_model: Feature dimension of the input.
101
+ """
102
+ super().__init__()
103
+ self.layers: nn.ModuleList = clones(decoder_layer, n)
104
+ self.layer_norm: LayerNorm = LayerNorm(d_model)
105
+
106
+ def forward(
107
+ self,
108
+ x: Tensor,
109
+ encoder_output: Tensor,
110
+ src_mask: Tensor,
111
+ tgt_mask: Tensor,
112
+ ) -> Tensor:
113
+ """Run ``x`` through every decoder layer.
114
+
115
+ Args:
116
+ x: Target input tensor of shape ``(batch, tgt_seq, d_model)``.
117
+ encoder_output: Encoder output of shape
118
+ ``(batch, src_seq, d_model)``.
119
+ src_mask: Source attention mask broadcastable to
120
+ ``(batch, 1, tgt_seq, src_seq)``.
121
+ tgt_mask: Target attention mask broadcastable to
122
+ ``(batch, 1, tgt_seq, tgt_seq)``.
123
+
124
+ Returns:
125
+ A tensor of shape ``(batch, tgt_seq, d_model)``.
126
+ """
127
+ for layer in self.layers:
128
+ x = layer(x, encoder_output, src_mask, tgt_mask)
129
+ return cast("Tensor", self.layer_norm(x))
@@ -0,0 +1,39 @@
1
+ """Token embedding layer."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+ from typing import cast
7
+
8
+ from torch import Tensor, nn
9
+
10
+
11
+ class Embeddings(nn.Module):
12
+ """Embed token indices and scale them by ``sqrt(d_model)``.
13
+
14
+ Args:
15
+ d_model: Dimension of the embedding vectors.
16
+ vocab_size: Number of tokens in the vocabulary.
17
+ """
18
+
19
+ def __init__(self, d_model: int, vocab_size: int) -> None:
20
+ """Initialize the embedding table.
21
+
22
+ Args:
23
+ d_model: Dimension of the embedding vectors.
24
+ vocab_size: Number of tokens in the vocabulary.
25
+ """
26
+ super().__init__()
27
+ self.embeddings: nn.Embedding = nn.Embedding(vocab_size, d_model)
28
+ self.d_model: int = d_model
29
+
30
+ def forward(self, x: Tensor) -> Tensor:
31
+ """Embed ``x`` and scale by ``sqrt(d_model)``.
32
+
33
+ Args:
34
+ x: Token indices of shape ``(batch, seq)``.
35
+
36
+ Returns:
37
+ Embeddings of shape ``(batch, seq, d_model)``.
38
+ """
39
+ return cast("Tensor", self.embeddings(x) * math.sqrt(self.d_model))
@@ -0,0 +1,101 @@
1
+ """Transformer encoder layers and stack."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import TYPE_CHECKING, cast
6
+
7
+ from torch import Tensor, nn
8
+
9
+ from .utils import LayerNorm, SublayerConnection, clones
10
+
11
+ if TYPE_CHECKING:
12
+ from .feed_forward import FeedForward
13
+ from .multi_head_attention import MultiHeadedAttention
14
+
15
+
16
+ class EncoderLayer(nn.Module):
17
+ """A single encoder layer with self-attention and a feed-forward network.
18
+
19
+ Args:
20
+ d_model: Feature dimension of the input.
21
+ self_attn: Multi-head self-attention module.
22
+ ff: Position-wise feed-forward network.
23
+ dropout: Dropout probability used inside the sublayer connections.
24
+ """
25
+
26
+ def __init__(
27
+ self,
28
+ d_model: int,
29
+ self_attn: MultiHeadedAttention,
30
+ ff: FeedForward,
31
+ dropout: float = 0.1,
32
+ ) -> None:
33
+ """Initialize the attention and feed-forward sublayers.
34
+
35
+ Args:
36
+ d_model: Feature dimension of the input.
37
+ self_attn: Multi-head self-attention module.
38
+ ff: Position-wise feed-forward network.
39
+ dropout: Dropout probability used inside the sublayer connections.
40
+ """
41
+ super().__init__()
42
+ self.self_attn: MultiHeadedAttention = self_attn
43
+ self.ff: FeedForward = ff
44
+ self.sublayers: nn.ModuleList = nn.ModuleList(
45
+ [
46
+ SublayerConnection(d_model, dropout),
47
+ SublayerConnection(d_model, dropout),
48
+ ]
49
+ )
50
+
51
+ def forward(self, x: Tensor, mask: Tensor) -> Tensor:
52
+ """Apply self-attention then the feed-forward network.
53
+
54
+ Args:
55
+ x: Input tensor of shape ``(batch, seq, d_model)``.
56
+ mask: Attention mask broadcastable to
57
+ ``(batch, 1, seq, seq)``.
58
+
59
+ Returns:
60
+ A tensor of shape ``(batch, seq, d_model)``.
61
+ """
62
+ x = self.sublayers[0](x, lambda t: self.self_attn(t, t, t, mask))
63
+ x = self.sublayers[1](x, self.ff)
64
+ return x
65
+
66
+
67
+ class Encoder(nn.Module):
68
+ """Stack of ``N`` encoder layers followed by layer normalization.
69
+
70
+ Args:
71
+ encoder_layer: The encoder layer to clone.
72
+ n: Number of encoder layers in the stack.
73
+ d_model: Feature dimension of the input.
74
+ """
75
+
76
+ def __init__(self, encoder_layer: EncoderLayer, n: int, d_model: int) -> None:
77
+ """Initialize the cloned layers and final layer norm.
78
+
79
+ Args:
80
+ encoder_layer: The encoder layer to clone.
81
+ n: Number of encoder layers in the stack.
82
+ d_model: Feature dimension of the input.
83
+ """
84
+ super().__init__()
85
+ self.layers: nn.ModuleList = clones(encoder_layer, n)
86
+ self.layer_norm: LayerNorm = LayerNorm(d_model)
87
+
88
+ def forward(self, x: Tensor, mask: Tensor) -> Tensor:
89
+ """Run ``x`` through every encoder layer.
90
+
91
+ Args:
92
+ x: Input tensor of shape ``(batch, seq, d_model)``.
93
+ mask: Attention mask broadcastable to
94
+ ``(batch, 1, seq, seq)``.
95
+
96
+ Returns:
97
+ A tensor of shape ``(batch, seq, d_model)``.
98
+ """
99
+ for layer in self.layers:
100
+ x = layer(x, mask)
101
+ return cast("Tensor", self.layer_norm(x))
@@ -0,0 +1,104 @@
1
+ """Full encoder-decoder transformer model."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import TYPE_CHECKING, cast
6
+
7
+ from torch import Tensor, nn
8
+
9
+ if TYPE_CHECKING:
10
+ from .decoder import Decoder
11
+ from .embeddings import Embeddings
12
+ from .encoder import Encoder
13
+ from .utils import Generator
14
+
15
+
16
+ class EncoderDecoder(nn.Module):
17
+ """Combine an encoder, decoder, generator and token embeddings.
18
+
19
+ Args:
20
+ encoder: The encoder stack.
21
+ decoder: The decoder stack.
22
+ generator: Output projection producing vocab log-probabilities.
23
+ src_emb: Source token embeddings.
24
+ tgt_emb: Target token embeddings.
25
+ """
26
+
27
+ def __init__(
28
+ self,
29
+ encoder: Encoder,
30
+ decoder: Decoder,
31
+ generator: Generator,
32
+ src_emb: Embeddings,
33
+ tgt_emb: Embeddings,
34
+ ) -> None:
35
+ """Store the encoder, decoder, generator and embeddings.
36
+
37
+ Args:
38
+ encoder: The encoder stack.
39
+ decoder: The decoder stack.
40
+ generator: Output projection producing vocab log-probabilities.
41
+ src_emb: Source token embeddings.
42
+ tgt_emb: Target token embeddings.
43
+ """
44
+ super().__init__()
45
+ self.encoder: Encoder = encoder
46
+ self.decoder: Decoder = decoder
47
+ self.generator: Generator = generator
48
+ self.src_emb: Embeddings = src_emb
49
+ self.tgt_emb: Embeddings = tgt_emb
50
+
51
+ def forward(
52
+ self,
53
+ src: Tensor,
54
+ target: Tensor,
55
+ src_mask: Tensor,
56
+ tgt_mask: Tensor,
57
+ ) -> Tensor:
58
+ """Encode ``src`` and decode ``target`` into vocab log-probabilities.
59
+
60
+ Args:
61
+ src: Source token indices of shape ``(batch, src_seq)``.
62
+ target: Target token indices of shape ``(batch, tgt_seq)``.
63
+ src_mask: Source attention mask.
64
+ tgt_mask: Target attention mask.
65
+
66
+ Returns:
67
+ Log-probabilities of shape ``(batch, tgt_seq, tgt_vocab_size)``.
68
+ """
69
+ memory = self.encode(src, src_mask)
70
+ decoder_output = self.decode(target, memory, src_mask, tgt_mask)
71
+ return cast("Tensor", self.generator(decoder_output))
72
+
73
+ def encode(self, src: Tensor, src_mask: Tensor) -> Tensor:
74
+ """Embed and encode the source tokens.
75
+
76
+ Args:
77
+ src: Source token indices of shape ``(batch, src_seq)``.
78
+ src_mask: Source attention mask.
79
+
80
+ Returns:
81
+ Encoder memory of shape ``(batch, src_seq, d_model)``.
82
+ """
83
+ return cast("Tensor", self.encoder(self.src_emb(src), src_mask))
84
+
85
+ def decode(
86
+ self,
87
+ tgt: Tensor,
88
+ encoder_output: Tensor,
89
+ src_mask: Tensor,
90
+ tgt_mask: Tensor,
91
+ ) -> Tensor:
92
+ """Embed and decode the target tokens given the encoder output.
93
+
94
+ Args:
95
+ tgt: Target token indices of shape ``(batch, tgt_seq)``.
96
+ encoder_output: Encoder memory of shape
97
+ ``(batch, src_seq, d_model)``.
98
+ src_mask: Source attention mask.
99
+ tgt_mask: Target attention mask.
100
+
101
+ Returns:
102
+ Decoder features of shape ``(batch, tgt_seq, d_model)``.
103
+ """
104
+ return cast("Tensor", self.decoder(self.tgt_emb(tgt), encoder_output, src_mask, tgt_mask))
@@ -0,0 +1,42 @@
1
+ """Position-wise feed-forward network."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import cast
6
+
7
+ from torch import Tensor, nn
8
+
9
+
10
+ class FeedForward(nn.Module):
11
+ """Two-layer position-wise feed-forward network with GELU activation.
12
+
13
+ Args:
14
+ d_model: Input and output feature dimension.
15
+ d_expand: Hidden (expanded) feature dimension.
16
+ dropout: Dropout probability applied between the linear layers.
17
+ """
18
+
19
+ def __init__(self, d_model: int, d_expand: int, dropout: float = 0.1) -> None:
20
+ """Initialize the linear layers and activation.
21
+
22
+ Args:
23
+ d_model: Input and output feature dimension.
24
+ d_expand: Hidden (expanded) feature dimension.
25
+ dropout: Dropout probability applied between the linear layers.
26
+ """
27
+ super().__init__()
28
+ self.dropout: nn.Dropout = nn.Dropout(dropout)
29
+ self.gelu: nn.GELU = nn.GELU()
30
+ self.a_1: nn.Linear = nn.Linear(d_model, d_expand)
31
+ self.a_2: nn.Linear = nn.Linear(d_expand, d_model)
32
+
33
+ def forward(self, x: Tensor) -> Tensor:
34
+ """Apply the feed-forward network to ``x``.
35
+
36
+ Args:
37
+ x: Input tensor of shape ``(batch, seq, d_model)``.
38
+
39
+ Returns:
40
+ A tensor of shape ``(batch, seq, d_model)``.
41
+ """
42
+ return cast("Tensor", self.a_2(self.dropout(self.gelu(self.a_1(x)))))
@@ -0,0 +1,58 @@
1
+ """Factory for building a complete encoder-decoder transformer."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import copy
6
+
7
+ from torch import nn
8
+
9
+ from .decoder import Decoder, DecoderLayer
10
+ from .embeddings import Embeddings
11
+ from .encoder import Encoder, EncoderLayer
12
+ from .encoder_decoder import EncoderDecoder
13
+ from .feed_forward import FeedForward
14
+ from .multi_head_attention import MultiHeadedAttention
15
+ from .utils import Generator
16
+
17
+
18
+ def make_model(
19
+ src_vocab_size: int,
20
+ tgt_vocab_size: int,
21
+ max_context: int = 1024,
22
+ N: int = 6,
23
+ d_model: int = 512,
24
+ d_ff: int = 2048,
25
+ h: int = 8,
26
+ dropout: float = 0.1,
27
+ ) -> EncoderDecoder:
28
+ """Build an encoder-decoder transformer with Xavier initialization.
29
+
30
+ Args:
31
+ src_vocab_size: Size of the source vocabulary.
32
+ tgt_vocab_size: Size of the target vocabulary.
33
+ max_context: Maximum sequence length for the RoPE tables.
34
+ N: Number of encoder and decoder layers.
35
+ d_model: Feature dimension of the model.
36
+ d_ff: Hidden dimension of the feed-forward network.
37
+ h: Number of attention heads.
38
+ dropout: Dropout probability used throughout the model.
39
+
40
+ Returns:
41
+ An initialized :class:`EncoderDecoder` model.
42
+ """
43
+ attn = MultiHeadedAttention(h, d_model, max_context, dropout)
44
+ cross_attn = copy.deepcopy(attn)
45
+ ff = FeedForward(d_model, d_ff, dropout)
46
+ model = EncoderDecoder(
47
+ Encoder(EncoderLayer(d_model, attn, ff, dropout), N, d_model),
48
+ Decoder(DecoderLayer(d_model, attn, cross_attn, ff, dropout), N, d_model),
49
+ Generator(d_model, tgt_vocab_size),
50
+ Embeddings(d_model, src_vocab_size),
51
+ Embeddings(d_model, tgt_vocab_size),
52
+ )
53
+
54
+ for p in model.parameters():
55
+ if p.dim() > 1:
56
+ nn.init.xavier_uniform_(p)
57
+
58
+ return model
@@ -0,0 +1,146 @@
1
+ """Multi-head attention with rotary position embeddings (RoPE)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+ from typing import cast
7
+
8
+ import torch
9
+ from torch import Tensor, nn
10
+
11
+ from .utils import clones
12
+
13
+
14
+ def attention(
15
+ query: Tensor,
16
+ key: Tensor,
17
+ value: Tensor,
18
+ mask: Tensor | None,
19
+ dropout: nn.Dropout | None,
20
+ ) -> tuple[Tensor, Tensor]:
21
+ """Compute scaled dot-product attention.
22
+
23
+ Args:
24
+ query: Query tensor of shape ``(batch, heads, seq_q, d_k)``.
25
+ key: Key tensor of shape ``(batch, heads, seq_k, d_k)``.
26
+ value: Value tensor of shape ``(batch, heads, seq_k, d_v)``.
27
+ mask: Optional mask broadcastable to ``(batch, heads, seq_q, seq_k)``.
28
+ Positions where the mask equals zero are ignored.
29
+ dropout: Optional dropout applied to the attention weights.
30
+
31
+ Returns:
32
+ A tuple ``(output, attn)`` where ``output`` has shape
33
+ ``(batch, heads, seq_q, d_v)`` and ``attn`` has shape
34
+ ``(batch, heads, seq_q, seq_k)``.
35
+ """
36
+ d_k = query.size(-1)
37
+ scores = torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(d_k)
38
+ if mask is not None:
39
+ scores = scores.masked_fill(mask == 0, -1e9)
40
+ p_attn = scores.softmax(dim=-1)
41
+ if dropout is not None:
42
+ p_attn = dropout(p_attn)
43
+ return torch.matmul(p_attn, value), p_attn
44
+
45
+
46
+ class MultiHeadedAttention(nn.Module):
47
+ """Multi-head self/cross attention using rotary position embeddings.
48
+
49
+ Args:
50
+ h: Number of attention heads.
51
+ d_model: Feature dimension of the input.
52
+ max_context: Maximum sequence length supported by the RoPE tables.
53
+ dropout: Dropout probability applied to the attention weights.
54
+ """
55
+
56
+ def __init__(
57
+ self,
58
+ h: int,
59
+ d_model: int,
60
+ max_context: int,
61
+ dropout: float = 0.1,
62
+ ) -> None:
63
+ """Initialize projections, dropout, and the RoPE buffers.
64
+
65
+ Args:
66
+ h: Number of attention heads.
67
+ d_model: Feature dimension of the input.
68
+ max_context: Maximum sequence length supported by the RoPE tables.
69
+ dropout: Dropout probability applied to the attention weights.
70
+
71
+ Raises:
72
+ AssertionError: If ``d_model`` is not divisible by ``h``.
73
+ """
74
+ super().__init__()
75
+ assert d_model % h == 0
76
+ self.d_k: int = d_model // h
77
+ self.h: int = h
78
+ self.linears: nn.ModuleList = clones(nn.Linear(d_model, d_model), 4)
79
+ self.attn: Tensor | None = None
80
+ self.dropout: nn.Dropout = nn.Dropout(p=dropout)
81
+
82
+ y_pos = torch.arange(max_context)
83
+ freqs = 10000 ** (-2 * torch.arange(self.d_k // 2, dtype=torch.float) / self.d_k)
84
+ angles = torch.outer(y_pos, freqs)
85
+ self.cos_table: Tensor
86
+ self.sin_table: Tensor
87
+ self.register_buffer("cos_table", torch.cos(angles))
88
+ self.register_buffer("sin_table", torch.sin(angles))
89
+
90
+ def rope_rotate(self, tensor: Tensor) -> Tensor:
91
+ """Apply rotary position embeddings to ``tensor``.
92
+
93
+ Args:
94
+ tensor: Tensor of shape ``(batch, heads, seq, d_k)``.
95
+
96
+ Returns:
97
+ The rotated tensor with the same shape as ``tensor``.
98
+ """
99
+ input_len = tensor.shape[2]
100
+
101
+ cos = self.cos_table[:input_len].unsqueeze(0).unsqueeze(0)
102
+ sin = self.sin_table[:input_len].unsqueeze(0).unsqueeze(0)
103
+
104
+ first_half = tensor[..., : self.d_k // 2]
105
+ second_half = tensor[..., self.d_k // 2 :]
106
+
107
+ a = first_half * cos - second_half * sin
108
+ b = first_half * sin + second_half * cos
109
+ return torch.concat((a, b), -1)
110
+
111
+ def forward(
112
+ self,
113
+ query: Tensor,
114
+ key: Tensor,
115
+ value: Tensor,
116
+ mask: Tensor | None,
117
+ ) -> Tensor:
118
+ """Apply multi-head attention to ``query``, ``key`` and ``value``.
119
+
120
+ Args:
121
+ query: Query tensor of shape ``(batch, seq_q, d_model)``.
122
+ key: Key tensor of shape ``(batch, seq_k, d_model)``.
123
+ value: Value tensor of shape ``(batch, seq_k, d_model)``.
124
+ mask: Optional mask broadcastable to
125
+ ``(batch, 1, seq_q, seq_k)``. Positions where the mask equals
126
+ zero are ignored.
127
+
128
+ Returns:
129
+ The attention output of shape ``(batch, seq_q, d_model)``.
130
+ """
131
+ if mask is not None:
132
+ mask = mask.unsqueeze(1)
133
+ nbatches = query.size(0)
134
+
135
+ query, key, value = [
136
+ lin(x).view(nbatches, -1, self.h, self.d_k).transpose(1, 2)
137
+ for lin, x in zip(self.linears, (query, key, value), strict=False)
138
+ ]
139
+
140
+ rotated_query = self.rope_rotate(query)
141
+ rotated_key = self.rope_rotate(key)
142
+
143
+ x, self.attn = attention(rotated_query, rotated_key, value, mask=mask, dropout=self.dropout)
144
+
145
+ x = x.transpose(1, 2).contiguous().view(nbatches, -1, self.h * self.d_k)
146
+ return cast("Tensor", self.linears[-1](x))
File without changes
@@ -0,0 +1,126 @@
1
+ """Reusable building blocks for the encoder-decoder transformer."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import copy
6
+ from typing import TYPE_CHECKING, cast
7
+
8
+ import torch
9
+ from torch import Tensor, nn
10
+
11
+ if TYPE_CHECKING:
12
+ from collections.abc import Callable
13
+
14
+
15
+ def clones[ModuleT: nn.Module](module: ModuleT, n: int) -> nn.ModuleList:
16
+ """Create ``n`` independent deep copies of ``module``.
17
+
18
+ Args:
19
+ module: The module to duplicate.
20
+ n: Number of copies to create.
21
+
22
+ Returns:
23
+ A ``ModuleList`` holding the ``n`` deep copies of ``module``.
24
+ """
25
+ return nn.ModuleList([copy.deepcopy(module) for _ in range(n)])
26
+
27
+
28
+ class LayerNorm(nn.Module):
29
+ """Layer normalization over the last dimension of the input.
30
+
31
+ Args:
32
+ features: Size of the normalized (last) dimension.
33
+ eps: Small constant added to the standard deviation for numerical
34
+ stability.
35
+ """
36
+
37
+ def __init__(self, features: int, eps: float = 1e-6) -> None:
38
+ """Initialize the learnable affine parameters.
39
+
40
+ Args:
41
+ features: Size of the normalized (last) dimension.
42
+ eps: Small constant added to the standard deviation.
43
+ """
44
+ super().__init__()
45
+ self.a_2: nn.Parameter = nn.Parameter(torch.ones(features))
46
+ self.b_2: nn.Parameter = nn.Parameter(torch.zeros(features))
47
+ self.eps: float = eps
48
+
49
+ def forward(self, x: Tensor) -> Tensor:
50
+ """Normalize ``x`` along its last dimension.
51
+
52
+ Args:
53
+ x: Input tensor of shape ``(..., features)``.
54
+
55
+ Returns:
56
+ A tensor with the same shape as ``x``, normalized along the last
57
+ dimension.
58
+ """
59
+ mean = x.mean(-1, keepdim=True)
60
+ std = x.std(-1, keepdim=True)
61
+ return self.a_2 * (x - mean) / (std + self.eps) + self.b_2
62
+
63
+
64
+ class SublayerConnection(nn.Module):
65
+ """Residual connection wrapped around a pre-normalized sublayer.
66
+
67
+ Args:
68
+ d_model: Feature dimension of the input.
69
+ dropout: Dropout probability applied to the sublayer output.
70
+ """
71
+
72
+ def __init__(self, d_model: int, dropout: float = 0.1) -> None:
73
+ """Initialize the layer norm and dropout module.
74
+
75
+ Args:
76
+ d_model: Feature dimension of the input.
77
+ dropout: Dropout probability applied to the sublayer output.
78
+ """
79
+ super().__init__()
80
+ self.dropout: nn.Dropout = nn.Dropout(p=dropout)
81
+ self.layer_norm: LayerNorm = LayerNorm(d_model)
82
+
83
+ def forward(self, x: Tensor, layer: Callable[[Tensor], Tensor]) -> Tensor:
84
+ """Apply ``layer`` to the normalized input and add the residual.
85
+
86
+ Args:
87
+ x: Input tensor of shape ``(..., d_model)``.
88
+ layer: Callable applied to the normalized input.
89
+
90
+ Returns:
91
+ A tensor of shape ``(..., d_model)`` produced by
92
+ ``x + dropout(layer(layer_norm(x)))``.
93
+ """
94
+ x_norm = self.layer_norm(x)
95
+ output_x = layer(x_norm)
96
+ return cast("Tensor", x + self.dropout(output_x))
97
+
98
+
99
+ class Generator(nn.Module):
100
+ """Map decoder features to log-probabilities over a vocabulary.
101
+
102
+ Args:
103
+ d_model: Feature dimension of the decoder output.
104
+ vocab_size: Number of tokens in the target vocabulary.
105
+ """
106
+
107
+ def __init__(self, d_model: int, vocab_size: int) -> None:
108
+ """Initialize the output projection.
109
+
110
+ Args:
111
+ d_model: Feature dimension of the decoder output.
112
+ vocab_size: Number of tokens in the target vocabulary.
113
+ """
114
+ super().__init__()
115
+ self.vocab_map: nn.Linear = nn.Linear(d_model, vocab_size)
116
+
117
+ def forward(self, x: Tensor) -> Tensor:
118
+ """Project ``x`` and apply a log-softmax over the vocabulary.
119
+
120
+ Args:
121
+ x: Decoder features of shape ``(batch, seq, d_model)``.
122
+
123
+ Returns:
124
+ Log-probabilities of shape ``(batch, seq, vocab_size)``.
125
+ """
126
+ return torch.log_softmax(self.vocab_map(x), -1)
@@ -0,0 +1,85 @@
1
+ Metadata-Version: 2.4
2
+ Name: encoder-decoder-transformer
3
+ Version: 0.1.0
4
+ Summary: A from-scratch PyTorch implementation of the Transformer encoder-decoder architecture.
5
+ Keywords: transformer,pytorch,attention,encoder-decoder,deep-learning
6
+ Author: Joe Rice
7
+ Author-email: Joe Rice <2Joerice@gmail.com>
8
+ License-Expression: MIT
9
+ Classifier: Development Status :: 3 - Alpha
10
+ Classifier: Intended Audience :: Developers
11
+ Classifier: Intended Audience :: Science/Research
12
+ Classifier: Programming Language :: Python :: 3
13
+ Classifier: Programming Language :: Python :: 3.13
14
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
15
+ Requires-Dist: torch>=2.14.1
16
+ Requires-Python: >=3.13
17
+ Description-Content-Type: text/markdown
18
+
19
+ # encoder-decoder-transformer
20
+
21
+ A from-scratch PyTorch implementation of the Transformer encoder-decoder
22
+ architecture, including multi-head attention with rotary position embeddings
23
+ (RoPE), pre-layer normalization, and a position-wise feed-forward network.
24
+
25
+ ## Installation
26
+
27
+ Install from source with [uv](https://docs.astral.sh/uv/):
28
+
29
+ ```bash
30
+ uv add encoder-decoder-transformer
31
+ ```
32
+
33
+ or with pip:
34
+
35
+ ```bash
36
+ pip install encoder-decoder-transformer
37
+ ```
38
+
39
+ The package targets Python 3.13 and depends on PyTorch.
40
+
41
+ ## Usage
42
+
43
+ Build a model and run a forward pass. Token indices are integer tensors of
44
+ shape `(batch, seq)`; the model returns log-probabilities of shape
45
+ `(batch, tgt_seq, tgt_vocab_size)`. Attention masks are boolean/integer tensors
46
+ of shape `(batch, seq_q, seq_k)` (the model adds the head dimension).
47
+
48
+ ```python
49
+ import torch
50
+
51
+ from encoder_decoder import make_model
52
+
53
+ model = make_model(
54
+ src_vocab_size=1000,
55
+ tgt_vocab_size=1000,
56
+ max_context=128,
57
+ N=2,
58
+ d_model=64,
59
+ d_ff=256,
60
+ h=4,
61
+ dropout=0.1,
62
+ )
63
+
64
+ src = torch.randint(0, 1000, (2, 8))
65
+ tgt = torch.randint(0, 1000, (2, 8))
66
+ src_mask = torch.ones(2, 8, 8)
67
+ tgt_mask = torch.ones(2, 8, 8)
68
+
69
+ log_probs = model(src, tgt, src_mask, tgt_mask)
70
+ print(log_probs.shape) # torch.Size([2, 8, 1000])
71
+ ```
72
+
73
+ ## Development
74
+
75
+ Run the linters and type checker with:
76
+
77
+ ```bash
78
+ uv run ruff check .
79
+ uv run ruff format --check .
80
+ uv run mypy
81
+ ```
82
+
83
+ ## License
84
+
85
+ MIT
@@ -0,0 +1,13 @@
1
+ encoder_decoder_transformer/__init__.py,sha256=PyP_wymKYM9ltnNf4uObAF0vQNJxipdolNQLIm4c--c,725
2
+ encoder_decoder_transformer/decoder.py,sha256=-whNa1m6UyMc5R18IkHQqxMgfyfJNqU-SHo6zBS6VQ0,4360
3
+ encoder_decoder_transformer/embeddings.py,sha256=JcgOwaolD51ewZ7qKiwTkKqQkc5t3KfJwWRjWYF1e-w,1079
4
+ encoder_decoder_transformer/encoder.py,sha256=qPwr-dn2GiwJXfDNtNBWeBhRBzpxmuuKY13AEAYNl_o,3242
5
+ encoder_decoder_transformer/encoder_decoder.py,sha256=H-h7fqhVrJRyIY1bDFctY6TJzrmNeIuuvTYSd86-zeY,3288
6
+ encoder_decoder_transformer/feed_forward.py,sha256=7vztLtmg2Z6FMYGs_2eBVL6sfskmDfD6ONUp8XRv8eI,1377
7
+ encoder_decoder_transformer/make_model.py,sha256=U-cVQDMjUxDJWHAZ30N-M2Dj5XQj58CfX_2-86NdjFo,1828
8
+ encoder_decoder_transformer/multi_head_attention.py,sha256=kkDAexGpPKizRAsq3IUh9OQtYHulfHTzk85R0fklCU8,4958
9
+ encoder_decoder_transformer/py.typed,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
10
+ encoder_decoder_transformer/utils.py,sha256=_WGUivOegGokiE1JAflOowqlIzrtQGHHdABG31TYZ8A,4001
11
+ encoder_decoder_transformer-0.1.0.dist-info/WHEEL,sha256=cmC5s21ojypbVslldL7IJq3hjZH-tINy4rziKePFsG0,81
12
+ encoder_decoder_transformer-0.1.0.dist-info/METADATA,sha256=jePIVQ5u6KHdJfifuuwK-8Sm7CgcsONvOEVVDICLtFA,2143
13
+ encoder_decoder_transformer-0.1.0.dist-info/RECORD,,
@@ -0,0 +1,4 @@
1
+ Wheel-Version: 1.0
2
+ Generator: uv 0.12.23
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any