Views
No views yet
A complete implementation of the Transformer architecture from the paper Attention Is All You Need, built entirely with PyTorch.
torch.nn.Transformer.1flowchart TD
2
3A[Source Tokens]
4B[Embedding]
5C[Positional Encoding]
6
7D["Encoder × N"]
8
9E[Encoder Memory]
10
11F[Target Tokens]
12G[Embedding]
13H[Positional Encoding]
14
15I["Decoder × N"]
16
17J[Linear Layer]
18
19K[Vocabulary Probabilities]
20
21A --> B --> C --> D --> E
22
23F --> G --> H --> I
24
25E --> I
26
27I --> J --> K1graph TD
2
3Transformer
4
5Transformer --> Embedding
6Transformer --> PositionalEncoding
7Transformer --> Encoder
8Transformer --> Decoder
9Transformer --> Linear
10
11Encoder --> MultiHeadAttention
12Encoder --> FeedForward
13Encoder --> LayerNorm
14
15Decoder --> MaskedAttention
16Decoder --> CrossAttention
17Decoder --> FeedForward2
18Decoder --> LayerNorm21Input
2 │
3 ▼
4Multi-Head Self Attention
5 │
6Add & LayerNorm
7 │
8Feed Forward Network
9 │
10Add & LayerNorm
11 │
12Output1Input
2 │
3 ▼
4Masked Multi-Head Attention
5 │
6Add & LayerNorm
7 │
8Cross Attention
9 │
10Add & LayerNorm
11 │
12Feed Forward Network
13 │
14Add & LayerNorm
15 │
16Output1transformer-from-scratch/
2
3├── model.py
4├── encoder.py
5├── decoder.py
6├── attention.py
7├── positional_encoding.py
8├── config.py
9├── train.py
10├── inference.py
11├── README.md
12│
13└── notebooks/| Hyperparameter | Value |
|---|---|
| Encoder Layers | 6 |
| Decoder Layers | 6 |
| Attention Heads | 8 |
| Embedding Size | 512 |
| Feed Forward Size | 2048 |
| Maximum Sequence Length | 5000 |
1import torch
2from model import Transformer
3
4src = torch.randint(0, 10000, (64, 20))
5tgt = torch.randint(0, 12000, (64, 15))
6
7model = Transformer(
8 src_vocab_size=10000,
9 tgt_vocab_size=12000,
10 num_heads=8,
11 num_layers=6,
12 emb_dim=512,
13 nn_dim=2048
14)
15
16output = model(src, tgt)
17
18print(output.shape)torch.Size([64, 15, 12000])1sequenceDiagram
2
3participant Source
4participant Encoder
5participant Decoder
6participant Output
7
8Source->>Encoder: Source Tokens
9
10Encoder->>Encoder: Self Attention
11
12Encoder-->>Decoder: Encoder Memory
13
14Decoder->>Decoder: Masked Self Attention
15
16Decoder->>Encoder: Cross Attention
17
18Decoder->>Output: Vocabulary Logits1criterion = nn.CrossEntropyLoss()
2
3optimizer = torch.optim.Adam(
4 model.parameters(),
5 lr=1e-4
6)