Upload folder using huggingface_hub
Browse files- README.md +138 -0
- config.json +21 -0
- model.py +290 -0
- pytorch_model.pt +3 -0
- tokenizer_config.json +7 -0
README.md
CHANGED
|
@@ -1,3 +1,141 @@
|
|
| 1 |
---
|
|
|
|
| 2 |
license: mit
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 3 |
---
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
---
|
| 2 |
+
language: en
|
| 3 |
license: mit
|
| 4 |
+
tags:
|
| 5 |
+
- pytorch
|
| 6 |
+
- language-model
|
| 7 |
+
- transformer
|
| 8 |
+
- decoder-only
|
| 9 |
+
- custom-architecture
|
| 10 |
+
- text-generation
|
| 11 |
+
- educational
|
| 12 |
+
- from-scratch
|
| 13 |
+
datasets:
|
| 14 |
+
- Skylion007/openwebtext
|
| 15 |
+
library_name: pytorch
|
| 16 |
---
|
| 17 |
+
|
| 18 |
+
# First5M — Decoder-Only Transformer Language Model
|
| 19 |
+
|
| 20 |
+
A ~5M parameter GPT-style language model built **entirely from scratch** using PyTorch.
|
| 21 |
+
Trained on [OpenWebText](https://huggingface.co/datasets/Skylion007/openwebtext)
|
| 22 |
+
following the architecture from ["Attention Is All You Need"](https://arxiv.org/abs/1706.03762)
|
| 23 |
+
(Vaswani et al., 2017).
|
| 24 |
+
|
| 25 |
+
This is an **educational project** — the goal is to understand every line of code
|
| 26 |
+
in a Transformer, not to build a production model.
|
| 27 |
+
|
| 28 |
+
## Model Details
|
| 29 |
+
|
| 30 |
+
| Property | Value |
|
| 31 |
+
|---|---|
|
| 32 |
+
| Architecture | Decoder-only Transformer (Pre-LayerNorm) |
|
| 33 |
+
| Parameters | ~5M |
|
| 34 |
+
| Layers | 6 |
|
| 35 |
+
| Hidden Dimension (d_model) | 256 |
|
| 36 |
+
| Attention Heads | 4 |
|
| 37 |
+
| FFN Dimension (d_ff) | 1024 |
|
| 38 |
+
| Context Window | 256 tokens |
|
| 39 |
+
| Tokenizer | tiktoken GPT-2 BPE (50,257 vocab) |
|
| 40 |
+
| Training Data | OpenWebText (~328M tokens, 17999 steps) |
|
| 41 |
+
| Best Val Loss | 4.986894807815552 (PPL 146) |
|
| 42 |
+
| Positional Encoding | Sinusoidal (not learned) |
|
| 43 |
+
| Weight Tying | Yes (embedding = output head) |
|
| 44 |
+
|
| 45 |
+
## Quick Start
|
| 46 |
+
|
| 47 |
+
**Requirements:** `pip install torch tiktoken huggingface_hub`
|
| 48 |
+
|
| 49 |
+
```python
|
| 50 |
+
import torch
|
| 51 |
+
import tiktoken
|
| 52 |
+
from huggingface_hub import hf_hub_download
|
| 53 |
+
import importlib.util
|
| 54 |
+
|
| 55 |
+
# Download model files
|
| 56 |
+
model_py_path = hf_hub_download("kaafivikrant/First5M", "model.py")
|
| 57 |
+
weights_path = hf_hub_download("kaafivikrant/First5M", "pytorch_model.pt")
|
| 58 |
+
|
| 59 |
+
# Load the model class from model.py
|
| 60 |
+
spec = importlib.util.spec_from_file_location("model", model_py_path)
|
| 61 |
+
mod = importlib.util.module_from_spec(spec)
|
| 62 |
+
spec.loader.exec_module(mod)
|
| 63 |
+
|
| 64 |
+
# Build model and load weights
|
| 65 |
+
config = mod.ModelConfig()
|
| 66 |
+
model = mod.TransformerLM(config)
|
| 67 |
+
state_dict = torch.load(weights_path, map_location="cpu", weights_only=True)
|
| 68 |
+
model.load_state_dict(state_dict, strict=False) # strict=False: lm_head is weight-tied
|
| 69 |
+
model.eval()
|
| 70 |
+
|
| 71 |
+
# Generate text
|
| 72 |
+
enc = tiktoken.get_encoding("gpt2")
|
| 73 |
+
prompt = "The meaning of life is"
|
| 74 |
+
ids = torch.tensor([enc.encode(prompt)], dtype=torch.long)
|
| 75 |
+
out = model.generate(ids, max_new_tokens=100, temperature=0.8, top_k=50)
|
| 76 |
+
print(enc.decode(out[0].tolist()))
|
| 77 |
+
```
|
| 78 |
+
|
| 79 |
+
## Architecture
|
| 80 |
+
|
| 81 |
+
```
|
| 82 |
+
Input token IDs [batch, seq_len]
|
| 83 |
+
|
|
| 84 |
+
Token Embedding [50257, 256]
|
| 85 |
+
+
|
| 86 |
+
Sinusoidal Positional Encoding
|
| 87 |
+
|
|
| 88 |
+
6x Transformer Blocks:
|
| 89 |
+
|-- LayerNorm
|
| 90 |
+
|-- Multi-Head Self-Attention (4 heads x 64 dims, causal mask)
|
| 91 |
+
|-- Residual Add
|
| 92 |
+
|-- LayerNorm
|
| 93 |
+
|-- Feed-Forward (256 -> 1024 -> 256, GELU)
|
| 94 |
+
|-- Residual Add
|
| 95 |
+
|
|
| 96 |
+
Final LayerNorm
|
| 97 |
+
|
|
| 98 |
+
Output Head [256, 50257] (tied with embedding)
|
| 99 |
+
|
|
| 100 |
+
Logits [batch, seq_len, 50257]
|
| 101 |
+
```
|
| 102 |
+
|
| 103 |
+
## Training Details
|
| 104 |
+
|
| 105 |
+
- **Dataset:** OpenWebText (Skylion007/openwebtext) — ~4.3B tokens total
|
| 106 |
+
- **Tokens Seen:** ~328M (~7.6% of dataset)
|
| 107 |
+
- **Optimizer:** AdamW (betas=0.9/0.95, weight_decay=0.1)
|
| 108 |
+
- **LR Schedule:** Cosine decay with linear warmup (500 steps)
|
| 109 |
+
- **Peak LR:** 3e-4, Min LR: 3e-5
|
| 110 |
+
- **Batch Size:** 64 effective (16 x 4 gradient accumulation steps)
|
| 111 |
+
- **Hardware:** Apple M1, 16GB RAM, MPS backend
|
| 112 |
+
- **Training Time:** ~22 hours
|
| 113 |
+
|
| 114 |
+
## Generation Parameters
|
| 115 |
+
|
| 116 |
+
The `generate()` method supports:
|
| 117 |
+
- `temperature`: Controls randomness (0.7-0.9 recommended)
|
| 118 |
+
- `top_k`: Limits sampling to top K tokens (40-50 recommended)
|
| 119 |
+
- `repetition_penalty`: Penalizes repeated tokens (1.2 default, 1.0 = off)
|
| 120 |
+
|
| 121 |
+
## Limitations
|
| 122 |
+
|
| 123 |
+
This is a small educational model. It:
|
| 124 |
+
- Produces low-quality, often incoherent text (expected for 5M params)
|
| 125 |
+
- Has a tiny context window (256 tokens)
|
| 126 |
+
- Has NOT been instruction-tuned or aligned
|
| 127 |
+
- May produce repetitive, nonsensical, or inappropriate text
|
| 128 |
+
- Is NOT intended for any production use
|
| 129 |
+
|
| 130 |
+
## Files
|
| 131 |
+
|
| 132 |
+
| File | Description |
|
| 133 |
+
|---|---|
|
| 134 |
+
| `model.py` | Standalone model class definitions (no dependencies beyond PyTorch) |
|
| 135 |
+
| `pytorch_model.pt` | Model weights (state dict) |
|
| 136 |
+
| `config.json` | Architecture hyperparameters |
|
| 137 |
+
| `tokenizer_config.json` | Tokenizer info (tiktoken GPT-2 BPE) |
|
| 138 |
+
|
| 139 |
+
## License
|
| 140 |
+
|
| 141 |
+
MIT
|
config.json
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architecture": "TransformerLM",
|
| 3 |
+
"model_type": "custom-decoder-only-transformer",
|
| 4 |
+
"tokenizer": "tiktoken-gpt2-bpe",
|
| 5 |
+
"weight_tying": true,
|
| 6 |
+
"pre_layernorm": true,
|
| 7 |
+
"positional_encoding": "sinusoidal",
|
| 8 |
+
"activation": "gelu",
|
| 9 |
+
"vocab_size": 50257,
|
| 10 |
+
"d_model": 256,
|
| 11 |
+
"n_heads": 4,
|
| 12 |
+
"n_layers": 6,
|
| 13 |
+
"d_ff": 1024,
|
| 14 |
+
"max_seq_len": 256,
|
| 15 |
+
"dropout": 0.1,
|
| 16 |
+
"training": {
|
| 17 |
+
"dataset": "Skylion007/openwebtext",
|
| 18 |
+
"steps": 17999,
|
| 19 |
+
"val_loss": 4.986894807815552
|
| 20 |
+
}
|
| 21 |
+
}
|
model.py
ADDED
|
@@ -0,0 +1,290 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
model.py - Standalone Transformer LM for inference.
|
| 3 |
+
|
| 4 |
+
A ~5M parameter decoder-only Transformer language model trained on OpenWebText.
|
| 5 |
+
Built from scratch following "Attention Is All You Need" (Vaswani et al., 2017).
|
| 6 |
+
|
| 7 |
+
Usage:
|
| 8 |
+
import torch, tiktoken
|
| 9 |
+
from model import ModelConfig, TransformerLM
|
| 10 |
+
|
| 11 |
+
config = ModelConfig()
|
| 12 |
+
model = TransformerLM(config)
|
| 13 |
+
|
| 14 |
+
state_dict = torch.load("pytorch_model.pt", map_location="cpu", weights_only=True)
|
| 15 |
+
model.load_state_dict(state_dict, strict=False) # strict=False: lm_head is weight-tied
|
| 16 |
+
model.eval()
|
| 17 |
+
|
| 18 |
+
enc = tiktoken.get_encoding("gpt2")
|
| 19 |
+
ids = torch.tensor([enc.encode("Once upon a time")])
|
| 20 |
+
out = model.generate(ids, max_new_tokens=100, temperature=0.8, top_k=50)
|
| 21 |
+
print(enc.decode(out[0].tolist()))
|
| 22 |
+
"""
|
| 23 |
+
|
| 24 |
+
import math
|
| 25 |
+
from dataclasses import dataclass
|
| 26 |
+
from typing import List, Optional, Tuple
|
| 27 |
+
|
| 28 |
+
import torch
|
| 29 |
+
import torch.nn as nn
|
| 30 |
+
import torch.nn.functional as F
|
| 31 |
+
|
| 32 |
+
HAS_SDPA = hasattr(F, "scaled_dot_product_attention")
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
@dataclass
|
| 36 |
+
class ModelConfig:
|
| 37 |
+
"""Architecture hyperparameters."""
|
| 38 |
+
vocab_size: int = 50257 # GPT-2 BPE vocabulary size
|
| 39 |
+
d_model: int = 256 # Hidden dimension
|
| 40 |
+
n_heads: int = 4 # Number of attention heads
|
| 41 |
+
n_layers: int = 6 # Number of Transformer blocks
|
| 42 |
+
d_ff: int = 1024 # Feed-forward inner dimension (4 * d_model)
|
| 43 |
+
max_seq_len: int = 256 # Maximum sequence length (context window)
|
| 44 |
+
dropout: float = 0.1 # Dropout rate
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
class SinusoidalPositionalEncoding(nn.Module):
|
| 48 |
+
"""Sinusoidal Positional Encoding (Section 3.5 of the original paper)."""
|
| 49 |
+
|
| 50 |
+
def __init__(self, d_model: int, max_seq_len: int = 5000, dropout: float = 0.1):
|
| 51 |
+
super().__init__()
|
| 52 |
+
self.dropout = nn.Dropout(p=dropout)
|
| 53 |
+
|
| 54 |
+
pe = torch.zeros(max_seq_len, d_model)
|
| 55 |
+
position = torch.arange(0, max_seq_len, dtype=torch.float).unsqueeze(1)
|
| 56 |
+
div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))
|
| 57 |
+
|
| 58 |
+
pe[:, 0::2] = torch.sin(position * div_term)
|
| 59 |
+
pe[:, 1::2] = torch.cos(position * div_term)
|
| 60 |
+
self.register_buffer("pe", pe.unsqueeze(0))
|
| 61 |
+
|
| 62 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 63 |
+
x = x + self.pe[:, :x.size(1)]
|
| 64 |
+
return self.dropout(x)
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
class MultiHeadSelfAttention(nn.Module):
|
| 68 |
+
"""Multi-Head Self-Attention with causal masking and KV-cache support."""
|
| 69 |
+
|
| 70 |
+
def __init__(self, d_model: int, n_heads: int, max_seq_len: int = 512, dropout: float = 0.1):
|
| 71 |
+
super().__init__()
|
| 72 |
+
assert d_model % n_heads == 0
|
| 73 |
+
self.n_heads = n_heads
|
| 74 |
+
self.d_k = d_model // n_heads
|
| 75 |
+
self.dropout = dropout
|
| 76 |
+
|
| 77 |
+
self.W_q = nn.Linear(d_model, d_model, bias=False)
|
| 78 |
+
self.W_k = nn.Linear(d_model, d_model, bias=False)
|
| 79 |
+
self.W_v = nn.Linear(d_model, d_model, bias=False)
|
| 80 |
+
self.W_o = nn.Linear(d_model, d_model, bias=False)
|
| 81 |
+
|
| 82 |
+
self.attn_dropout = nn.Dropout(dropout)
|
| 83 |
+
self.resid_dropout = nn.Dropout(dropout)
|
| 84 |
+
|
| 85 |
+
if not HAS_SDPA:
|
| 86 |
+
self.register_buffer(
|
| 87 |
+
"causal_mask",
|
| 88 |
+
torch.tril(torch.ones(max_seq_len, max_seq_len)).view(1, 1, max_seq_len, max_seq_len),
|
| 89 |
+
)
|
| 90 |
+
|
| 91 |
+
def forward(self, x: torch.Tensor, kv_cache: Optional[Tuple[torch.Tensor, torch.Tensor]] = None):
|
| 92 |
+
B, T, C = x.shape
|
| 93 |
+
q = self.W_q(x).view(B, T, self.n_heads, self.d_k).transpose(1, 2)
|
| 94 |
+
k = self.W_k(x).view(B, T, self.n_heads, self.d_k).transpose(1, 2)
|
| 95 |
+
v = self.W_v(x).view(B, T, self.n_heads, self.d_k).transpose(1, 2)
|
| 96 |
+
|
| 97 |
+
new_cache = None
|
| 98 |
+
if kv_cache is not None:
|
| 99 |
+
k_prev, v_prev = kv_cache
|
| 100 |
+
k = torch.cat([k_prev, k], dim=2)
|
| 101 |
+
v = torch.cat([v_prev, v], dim=2)
|
| 102 |
+
new_cache = (k, v)
|
| 103 |
+
|
| 104 |
+
if HAS_SDPA:
|
| 105 |
+
out = F.scaled_dot_product_attention(
|
| 106 |
+
q, k, v,
|
| 107 |
+
is_causal=(kv_cache is None),
|
| 108 |
+
dropout_p=self.dropout if self.training else 0.0,
|
| 109 |
+
)
|
| 110 |
+
else:
|
| 111 |
+
S = k.size(2)
|
| 112 |
+
attn = (q @ k.transpose(-2, -1)) * (self.d_k ** -0.5)
|
| 113 |
+
if kv_cache is None:
|
| 114 |
+
attn = attn.masked_fill(self.causal_mask[:, :, :T, :T] == 0, float("-inf"))
|
| 115 |
+
attn = self.attn_dropout(F.softmax(attn, dim=-1))
|
| 116 |
+
out = attn @ v
|
| 117 |
+
|
| 118 |
+
out = out.transpose(1, 2).contiguous().view(B, T, C)
|
| 119 |
+
return self.resid_dropout(self.W_o(out)), new_cache
|
| 120 |
+
|
| 121 |
+
|
| 122 |
+
class FeedForward(nn.Module):
|
| 123 |
+
"""Position-wise Feed-Forward Network: expand -> GELU -> contract."""
|
| 124 |
+
|
| 125 |
+
def __init__(self, d_model: int, d_ff: int, dropout: float = 0.1):
|
| 126 |
+
super().__init__()
|
| 127 |
+
self.net = nn.Sequential(
|
| 128 |
+
nn.Linear(d_model, d_ff),
|
| 129 |
+
nn.GELU(),
|
| 130 |
+
nn.Linear(d_ff, d_model),
|
| 131 |
+
nn.Dropout(dropout),
|
| 132 |
+
)
|
| 133 |
+
|
| 134 |
+
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
| 135 |
+
return self.net(x)
|
| 136 |
+
|
| 137 |
+
|
| 138 |
+
class TransformerBlock(nn.Module):
|
| 139 |
+
"""Pre-LayerNorm Transformer decoder block (attention + FFN + residuals)."""
|
| 140 |
+
|
| 141 |
+
def __init__(self, d_model: int, n_heads: int, d_ff: int, max_seq_len: int = 512, dropout: float = 0.1):
|
| 142 |
+
super().__init__()
|
| 143 |
+
self.ln1 = nn.LayerNorm(d_model)
|
| 144 |
+
self.attn = MultiHeadSelfAttention(d_model, n_heads, max_seq_len, dropout)
|
| 145 |
+
self.ln2 = nn.LayerNorm(d_model)
|
| 146 |
+
self.ff = FeedForward(d_model, d_ff, dropout)
|
| 147 |
+
|
| 148 |
+
def forward(self, x: torch.Tensor, kv_cache=None):
|
| 149 |
+
attn_out, new_cache = self.attn(self.ln1(x), kv_cache=kv_cache)
|
| 150 |
+
x = x + attn_out
|
| 151 |
+
x = x + self.ff(self.ln2(x))
|
| 152 |
+
return x, new_cache
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
class TransformerLM(nn.Module):
|
| 156 |
+
"""
|
| 157 |
+
Decoder-only Transformer Language Model.
|
| 158 |
+
|
| 159 |
+
Features:
|
| 160 |
+
- Pre-LayerNorm architecture (GPT-2 style)
|
| 161 |
+
- Sinusoidal positional encoding
|
| 162 |
+
- Weight tying between embedding and output head
|
| 163 |
+
- KV-cache for efficient autoregressive generation
|
| 164 |
+
- Repetition penalty for better generation quality
|
| 165 |
+
"""
|
| 166 |
+
|
| 167 |
+
def __init__(self, config: ModelConfig):
|
| 168 |
+
super().__init__()
|
| 169 |
+
self.config = config
|
| 170 |
+
|
| 171 |
+
self.token_embedding = nn.Embedding(config.vocab_size, config.d_model)
|
| 172 |
+
self.pos_encoding = SinusoidalPositionalEncoding(
|
| 173 |
+
config.d_model, config.max_seq_len + 2048, config.dropout
|
| 174 |
+
)
|
| 175 |
+
self.blocks = nn.ModuleList([
|
| 176 |
+
TransformerBlock(config.d_model, config.n_heads, config.d_ff,
|
| 177 |
+
config.max_seq_len, config.dropout)
|
| 178 |
+
for _ in range(config.n_layers)
|
| 179 |
+
])
|
| 180 |
+
self.ln_f = nn.LayerNorm(config.d_model)
|
| 181 |
+
self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False)
|
| 182 |
+
|
| 183 |
+
# Weight tying: embedding and output head share the same weights
|
| 184 |
+
self.token_embedding.weight = self.lm_head.weight
|
| 185 |
+
|
| 186 |
+
self.apply(self._init_weights)
|
| 187 |
+
for pn, p in self.named_parameters():
|
| 188 |
+
if pn.endswith("W_o.weight") or pn.endswith("net.2.weight"):
|
| 189 |
+
torch.nn.init.normal_(p, mean=0.0, std=0.02 / math.sqrt(2 * config.n_layers))
|
| 190 |
+
|
| 191 |
+
def _init_weights(self, module):
|
| 192 |
+
if isinstance(module, nn.Linear):
|
| 193 |
+
torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
|
| 194 |
+
if module.bias is not None:
|
| 195 |
+
torch.nn.init.zeros_(module.bias)
|
| 196 |
+
elif isinstance(module, nn.Embedding):
|
| 197 |
+
torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)
|
| 198 |
+
|
| 199 |
+
def forward(self, idx, targets=None):
|
| 200 |
+
x = self.pos_encoding(self.token_embedding(idx))
|
| 201 |
+
for block in self.blocks:
|
| 202 |
+
x, _ = block(x)
|
| 203 |
+
logits = self.lm_head(self.ln_f(x))
|
| 204 |
+
loss = None
|
| 205 |
+
if targets is not None:
|
| 206 |
+
loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1))
|
| 207 |
+
return logits, loss
|
| 208 |
+
|
| 209 |
+
@torch.no_grad()
|
| 210 |
+
def generate(self, idx, max_new_tokens, temperature=1.0, top_k=None,
|
| 211 |
+
repetition_penalty=1.2):
|
| 212 |
+
"""
|
| 213 |
+
Autoregressive text generation with KV-cache and repetition penalty.
|
| 214 |
+
|
| 215 |
+
Args:
|
| 216 |
+
idx: Prompt token IDs, shape [batch, prompt_len]
|
| 217 |
+
max_new_tokens: Number of tokens to generate
|
| 218 |
+
temperature: Sampling temperature (default 1.0)
|
| 219 |
+
top_k: Only sample from top K tokens (default None = all)
|
| 220 |
+
repetition_penalty: Penalty for repeated tokens (1.0 = off, 1.2 = default)
|
| 221 |
+
|
| 222 |
+
Returns:
|
| 223 |
+
idx: Prompt + generated tokens, shape [batch, prompt_len + max_new_tokens]
|
| 224 |
+
"""
|
| 225 |
+
kv_caches: List[Optional[Tuple[torch.Tensor, torch.Tensor]]] = [None] * len(self.blocks)
|
| 226 |
+
|
| 227 |
+
# Phase 1: Prefill — process entire prompt
|
| 228 |
+
x = self.pos_encoding(self.token_embedding(idx))
|
| 229 |
+
for i, block in enumerate(self.blocks):
|
| 230 |
+
x, kv_caches[i] = block(x)
|
| 231 |
+
|
| 232 |
+
logits = self.lm_head(self.ln_f(x))
|
| 233 |
+
logits = logits[:, -1, :]
|
| 234 |
+
|
| 235 |
+
if repetition_penalty != 1.0:
|
| 236 |
+
for b in range(idx.size(0)):
|
| 237 |
+
seen = idx[b].unique()
|
| 238 |
+
for token_id in seen:
|
| 239 |
+
if logits[b, token_id] > 0:
|
| 240 |
+
logits[b, token_id] /= repetition_penalty
|
| 241 |
+
else:
|
| 242 |
+
logits[b, token_id] *= repetition_penalty
|
| 243 |
+
|
| 244 |
+
logits = logits / temperature
|
| 245 |
+
if top_k is not None:
|
| 246 |
+
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
|
| 247 |
+
logits[logits < v[:, [-1]]] = float("-inf")
|
| 248 |
+
|
| 249 |
+
probs = F.softmax(logits, dim=-1)
|
| 250 |
+
idx_next = torch.multinomial(probs, num_samples=1)
|
| 251 |
+
idx = torch.cat((idx, idx_next), dim=1)
|
| 252 |
+
|
| 253 |
+
# Phase 2: Decode — generate one token at a time with KV-cache
|
| 254 |
+
for _ in range(max_new_tokens - 1):
|
| 255 |
+
seq_pos = idx.size(1) - 1
|
| 256 |
+
x = self.token_embedding(idx_next)
|
| 257 |
+
x = x + self.pos_encoding.pe[:, seq_pos:seq_pos + 1]
|
| 258 |
+
|
| 259 |
+
for i, block in enumerate(self.blocks):
|
| 260 |
+
x, kv_caches[i] = block(x, kv_cache=kv_caches[i])
|
| 261 |
+
|
| 262 |
+
logits = self.lm_head(self.ln_f(x))
|
| 263 |
+
logits = logits[:, -1, :]
|
| 264 |
+
|
| 265 |
+
if repetition_penalty != 1.0:
|
| 266 |
+
for b in range(idx.size(0)):
|
| 267 |
+
seen = idx[b].unique()
|
| 268 |
+
for token_id in seen:
|
| 269 |
+
if logits[b, token_id] > 0:
|
| 270 |
+
logits[b, token_id] /= repetition_penalty
|
| 271 |
+
else:
|
| 272 |
+
logits[b, token_id] *= repetition_penalty
|
| 273 |
+
|
| 274 |
+
logits = logits / temperature
|
| 275 |
+
if top_k is not None:
|
| 276 |
+
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
|
| 277 |
+
logits[logits < v[:, [-1]]] = float("-inf")
|
| 278 |
+
|
| 279 |
+
probs = F.softmax(logits, dim=-1)
|
| 280 |
+
idx_next = torch.multinomial(probs, num_samples=1)
|
| 281 |
+
idx = torch.cat((idx, idx_next), dim=1)
|
| 282 |
+
|
| 283 |
+
if idx.size(1) > self.config.max_seq_len:
|
| 284 |
+
for i in range(len(kv_caches)):
|
| 285 |
+
if kv_caches[i] is not None:
|
| 286 |
+
k, v_tensor = kv_caches[i]
|
| 287 |
+
kv_caches[i] = (k[:, :, -self.config.max_seq_len:, :],
|
| 288 |
+
v_tensor[:, :, -self.config.max_seq_len:, :])
|
| 289 |
+
|
| 290 |
+
return idx
|
pytorch_model.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:a53096960358060006c1cc654f29a2e9418049e07e3cc95849b00cd240be29dd
|
| 3 |
+
size 72777963
|
tokenizer_config.json
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"tokenizer_type": "tiktoken",
|
| 3 |
+
"encoding": "gpt2",
|
| 4 |
+
"vocab_size": 50257,
|
| 5 |
+
"install": "pip install tiktoken",
|
| 6 |
+
"usage": "import tiktoken; enc = tiktoken.get_encoding('gpt2')"
|
| 7 |
+
}
|