mirror of
https://github.com/priyanshujain/torchmlx.git
synced 2026-10-02 11:07:13 +00:00
rewrite tinystories example
This commit is contained in:
1 parent
b0ef50eec2
commit
7fa8ae82ad
1 file changed
+228
-121
+228
-121
@@ -1,6 +1,7 @@
|
|||||||
import argparse
|
import argparse
|
||||||
import random
|
import random
|
||||||
import urllib.request
|
import urllib.request
|
||||||
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import torchmlx as torch
|
import torchmlx as torch
|
||||||
@@ -8,125 +9,6 @@ from torchmlx import nn, optim
|
|||||||
import torchmlx.nn.functional as F
|
import torchmlx.nn.functional as F
|
||||||
|
|
||||||
|
|
||||||
DATA_URL = "https://huggingface.co/datasets/roneneldan/TinyStories/resolve/main/TinyStoriesV2-GPT4-valid.txt"
|
|
||||||
DATA_PATH = Path(__file__).with_name("TinyStoriesV2-GPT4-valid.txt")
|
|
||||||
|
|
||||||
|
|
||||||
class Attention(nn.Module):
|
|
||||||
def __init__(self, width, heads):
|
|
||||||
super().__init__()
|
|
||||||
self.heads = heads
|
|
||||||
self.head_width = width // heads
|
|
||||||
self.qkv = nn.Linear(width, width * 3)
|
|
||||||
self.output = nn.Linear(width, width)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
batch, length, width = x.shape
|
|
||||||
q, k, v = torch.chunk(self.qkv(x), 3, dim=-1)
|
|
||||||
q = torch.transpose(
|
|
||||||
q.reshape(batch, length, self.heads, self.head_width), 1, 2
|
|
||||||
)
|
|
||||||
k = torch.transpose(
|
|
||||||
k.reshape(batch, length, self.heads, self.head_width), 1, 2
|
|
||||||
)
|
|
||||||
v = torch.transpose(
|
|
||||||
v.reshape(batch, length, self.heads, self.head_width), 1, 2
|
|
||||||
)
|
|
||||||
x = F.scaled_dot_product_attention(q, k, v, is_causal=True)
|
|
||||||
x = torch.transpose(x, 1, 2).reshape(batch, length, width)
|
|
||||||
return self.output(x)
|
|
||||||
|
|
||||||
|
|
||||||
class Block(nn.Module):
|
|
||||||
def __init__(self, width, heads):
|
|
||||||
super().__init__()
|
|
||||||
self.attention_norm = nn.LayerNorm(width)
|
|
||||||
self.attention = Attention(width, heads)
|
|
||||||
self.feed_forward_norm = nn.LayerNorm(width)
|
|
||||||
self.feed_forward = nn.Sequential(
|
|
||||||
nn.Linear(width, width * 4),
|
|
||||||
nn.GELU(),
|
|
||||||
nn.Linear(width * 4, width),
|
|
||||||
)
|
|
||||||
|
|
||||||
def forward(self, x):
|
|
||||||
x = x + self.attention(self.attention_norm(x))
|
|
||||||
return x + self.feed_forward(self.feed_forward_norm(x))
|
|
||||||
|
|
||||||
|
|
||||||
class TinyStoriesModel(nn.Module):
|
|
||||||
def __init__(self, vocab_size, context, width=64, heads=4, layers=2):
|
|
||||||
super().__init__()
|
|
||||||
self.context = context
|
|
||||||
self.token_embedding = nn.Embedding(vocab_size, width)
|
|
||||||
self.position_embedding = nn.Embedding(context, width)
|
|
||||||
self.blocks = nn.ModuleList([Block(width, heads) for _ in range(layers)])
|
|
||||||
self.norm = nn.LayerNorm(width)
|
|
||||||
self.output = nn.Linear(width, vocab_size)
|
|
||||||
|
|
||||||
def forward(self, tokens):
|
|
||||||
positions = torch.arange(tokens.shape[1], device=device)
|
|
||||||
x = self.token_embedding(tokens) + self.position_embedding(positions)
|
|
||||||
for block in self.blocks:
|
|
||||||
x = block(x)
|
|
||||||
return self.output(self.norm(x))
|
|
||||||
|
|
||||||
|
|
||||||
def download_data():
|
|
||||||
if not DATA_PATH.exists():
|
|
||||||
print(f"downloading TinyStories to {DATA_PATH}")
|
|
||||||
urllib.request.urlretrieve(DATA_URL, DATA_PATH)
|
|
||||||
return DATA_PATH.read_text(encoding="utf-8")
|
|
||||||
|
|
||||||
|
|
||||||
def batch(encoded, batch_size, context):
|
|
||||||
starts = [random.randrange(len(encoded) - context - 1) for _ in range(batch_size)]
|
|
||||||
x = [encoded[start : start + context] for start in starts]
|
|
||||||
y = [encoded[start + 1 : start + context + 1] for start in starts]
|
|
||||||
return torch.tensor(x, dtype=torch.int64, device=device), torch.tensor(
|
|
||||||
y, dtype=torch.int64, device=device
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def generate(model, prompt, encode, decode, length):
|
|
||||||
tokens = encode(prompt)
|
|
||||||
for _ in range(length):
|
|
||||||
x = torch.tensor([tokens[-model.context :]], dtype=torch.int64, device=device)
|
|
||||||
logits = model(x)[:, -1, :]
|
|
||||||
token = torch.categorical(logits).item()
|
|
||||||
tokens.append(token)
|
|
||||||
return decode(tokens)
|
|
||||||
|
|
||||||
|
|
||||||
def main():
|
|
||||||
parser = argparse.ArgumentParser()
|
|
||||||
parser.add_argument("steps", nargs="?", type=int, default=20)
|
|
||||||
parser.add_argument("--no-compile", action="store_true")
|
|
||||||
args = parser.parse_args()
|
|
||||||
text = download_data()
|
|
||||||
characters = sorted(set(text))
|
|
||||||
to_id = {character: index for index, character in enumerate(characters)}
|
|
||||||
encode = lambda value: [to_id[character] for character in value]
|
|
||||||
decode = lambda values: "".join(characters[int(value)] for value in values)
|
|
||||||
encoded = encode(text)
|
|
||||||
context = 64
|
|
||||||
model = TinyStoriesModel(len(characters), context).to(device)
|
|
||||||
optimizer = optim.AdamW(model.parameters(), lr=3e-4)
|
|
||||||
loss_fn = lambda logits, targets: F.cross_entropy(
|
|
||||||
logits.reshape(-1, len(characters)), targets.reshape(-1)
|
|
||||||
)
|
|
||||||
trainer = torch.Trainer(
|
|
||||||
model, optimizer, loss_fn, compile=not args.no_compile
|
|
||||||
)
|
|
||||||
interval = max(1, args.steps // 10)
|
|
||||||
for step in range(args.steps):
|
|
||||||
x, y = batch(encoded, 8, context)
|
|
||||||
loss = trainer.step(x, y)
|
|
||||||
if step % interval == 0 or step == args.steps - 1:
|
|
||||||
print(f"step {step + 1}: loss {loss.item():.4f}")
|
|
||||||
print(generate(model, "Once upon a time", encode, decode, 120))
|
|
||||||
|
|
||||||
|
|
||||||
device = torch.device(
|
device = torch.device(
|
||||||
"mps"
|
"mps"
|
||||||
if torch.backends.mps.is_available()
|
if torch.backends.mps.is_available()
|
||||||
@@ -135,6 +17,231 @@ device = torch.device(
|
|||||||
else "cpu"
|
else "cpu"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
DATA_URL = "https://huggingface.co/datasets/roneneldan/TinyStories/resolve/main/TinyStoriesV2-GPT4-valid.txt"
|
||||||
|
DATA_PATH = Path(__file__).with_name("TinyStoriesV2-GPT4-valid.txt")
|
||||||
|
|
||||||
if __name__ == "__main__":
|
|
||||||
main()
|
@dataclass
|
||||||
|
class TransformerConfig:
|
||||||
|
vocabulary_size: int
|
||||||
|
context_length: int = 64
|
||||||
|
model_dimension: int = 64
|
||||||
|
head_count: int = 4
|
||||||
|
layer_count: int = 2
|
||||||
|
feed_forward_dimension: int = 256
|
||||||
|
|
||||||
|
|
||||||
|
class CharacterTokenizer:
|
||||||
|
def __init__(self, text):
|
||||||
|
self.characters = sorted(set(text))
|
||||||
|
self.token_by_character = {
|
||||||
|
character: token for token, character in enumerate(self.characters)
|
||||||
|
}
|
||||||
|
|
||||||
|
def encode(self, text):
|
||||||
|
return [self.token_by_character[character] for character in text]
|
||||||
|
|
||||||
|
def decode(self, tokens):
|
||||||
|
return "".join(self.characters[int(token)] for token in tokens)
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return len(self.characters)
|
||||||
|
|
||||||
|
|
||||||
|
class TokenEmbedding(nn.Module):
|
||||||
|
def __init__(self, config):
|
||||||
|
super().__init__()
|
||||||
|
self.embedding = nn.Embedding(
|
||||||
|
config.vocabulary_size, config.model_dimension
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, tokens):
|
||||||
|
return self.embedding(tokens)
|
||||||
|
|
||||||
|
|
||||||
|
class PositionalEmbedding(nn.Module):
|
||||||
|
def __init__(self, config):
|
||||||
|
super().__init__()
|
||||||
|
self.embedding = nn.Embedding(
|
||||||
|
config.context_length, config.model_dimension
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, tokens):
|
||||||
|
positions = torch.arange(tokens.shape[1], device=tokens.device)
|
||||||
|
return self.embedding(positions)
|
||||||
|
|
||||||
|
|
||||||
|
class CausalSelfAttention(nn.Module):
|
||||||
|
def __init__(self, config):
|
||||||
|
super().__init__()
|
||||||
|
self.head_count = config.head_count
|
||||||
|
self.head_dimension = config.model_dimension // config.head_count
|
||||||
|
self.query = nn.Linear(config.model_dimension, config.model_dimension)
|
||||||
|
self.key = nn.Linear(config.model_dimension, config.model_dimension)
|
||||||
|
self.value = nn.Linear(config.model_dimension, config.model_dimension)
|
||||||
|
self.output = nn.Linear(config.model_dimension, config.model_dimension)
|
||||||
|
|
||||||
|
def split_heads(self, hidden_states):
|
||||||
|
batch_size, sequence_length, _ = hidden_states.shape
|
||||||
|
hidden_states = hidden_states.reshape(
|
||||||
|
batch_size, sequence_length, self.head_count, self.head_dimension
|
||||||
|
)
|
||||||
|
return hidden_states.transpose(1, 2)
|
||||||
|
|
||||||
|
def merge_heads(self, hidden_states):
|
||||||
|
batch_size, _, sequence_length, _ = hidden_states.shape
|
||||||
|
hidden_states = hidden_states.transpose(1, 2).contiguous()
|
||||||
|
return hidden_states.reshape(
|
||||||
|
batch_size, sequence_length, self.head_count * self.head_dimension
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, hidden_states):
|
||||||
|
queries = self.split_heads(self.query(hidden_states))
|
||||||
|
keys = self.split_heads(self.key(hidden_states))
|
||||||
|
values = self.split_heads(self.value(hidden_states))
|
||||||
|
|
||||||
|
attention_scores = torch.matmul(queries, keys.transpose(-2, -1))
|
||||||
|
attention_scores = attention_scores / self.head_dimension**0.5
|
||||||
|
|
||||||
|
sequence_length = hidden_states.shape[1]
|
||||||
|
future_positions = torch.triu(
|
||||||
|
torch.ones(
|
||||||
|
sequence_length,
|
||||||
|
sequence_length,
|
||||||
|
dtype=torch.bool,
|
||||||
|
device=hidden_states.device,
|
||||||
|
),
|
||||||
|
diagonal=1,
|
||||||
|
)
|
||||||
|
attention_scores = attention_scores.masked_fill(
|
||||||
|
future_positions, float("-inf")
|
||||||
|
)
|
||||||
|
attention_weights = F.softmax(attention_scores, dim=-1)
|
||||||
|
attended_values = torch.matmul(attention_weights, values)
|
||||||
|
return self.output(self.merge_heads(attended_values))
|
||||||
|
|
||||||
|
|
||||||
|
class FeedForward(nn.Module):
|
||||||
|
def __init__(self, config):
|
||||||
|
super().__init__()
|
||||||
|
self.layers = nn.Sequential(
|
||||||
|
nn.Linear(config.model_dimension, config.feed_forward_dimension),
|
||||||
|
nn.GELU(),
|
||||||
|
nn.Linear(config.feed_forward_dimension, config.model_dimension),
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, hidden_states):
|
||||||
|
return self.layers(hidden_states)
|
||||||
|
|
||||||
|
|
||||||
|
class Unembedding(nn.Module):
|
||||||
|
def __init__(self, config):
|
||||||
|
super().__init__()
|
||||||
|
self.output = nn.Linear(
|
||||||
|
config.model_dimension, config.vocabulary_size, bias=False
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, hidden_states):
|
||||||
|
return self.output(hidden_states)
|
||||||
|
|
||||||
|
|
||||||
|
class TransformerBlock(nn.Module):
|
||||||
|
def __init__(self, config):
|
||||||
|
super().__init__()
|
||||||
|
self.attention_norm = nn.LayerNorm(config.model_dimension)
|
||||||
|
self.attention = CausalSelfAttention(config)
|
||||||
|
self.feed_forward_norm = nn.LayerNorm(config.model_dimension)
|
||||||
|
self.feed_forward = FeedForward(config)
|
||||||
|
|
||||||
|
def forward(self, hidden_states):
|
||||||
|
hidden_states = hidden_states + self.attention(
|
||||||
|
self.attention_norm(hidden_states)
|
||||||
|
)
|
||||||
|
return hidden_states + self.feed_forward(
|
||||||
|
self.feed_forward_norm(hidden_states)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TinyStoriesTransformer(nn.Module):
|
||||||
|
def __init__(self, config):
|
||||||
|
super().__init__()
|
||||||
|
self.config = config
|
||||||
|
self.token_embedding = TokenEmbedding(config)
|
||||||
|
self.position_embedding = PositionalEmbedding(config)
|
||||||
|
self.blocks = nn.ModuleList(
|
||||||
|
[TransformerBlock(config) for _ in range(config.layer_count)]
|
||||||
|
)
|
||||||
|
self.final_norm = nn.LayerNorm(config.model_dimension)
|
||||||
|
self.unembedding = Unembedding(config)
|
||||||
|
|
||||||
|
def forward(self, tokens):
|
||||||
|
hidden_states = self.token_embedding(tokens)
|
||||||
|
hidden_states = hidden_states + self.position_embedding(tokens)
|
||||||
|
for block in self.blocks:
|
||||||
|
hidden_states = block(hidden_states)
|
||||||
|
return self.unembedding(self.final_norm(hidden_states))
|
||||||
|
|
||||||
|
|
||||||
|
def load_tiny_stories():
|
||||||
|
if not DATA_PATH.exists():
|
||||||
|
print(f"downloading TinyStories to {DATA_PATH}")
|
||||||
|
urllib.request.urlretrieve(DATA_URL, DATA_PATH)
|
||||||
|
return DATA_PATH.read_text(encoding="utf-8")
|
||||||
|
|
||||||
|
|
||||||
|
def create_batch(encoded_text, batch_size, context_length):
|
||||||
|
starts = [
|
||||||
|
random.randrange(len(encoded_text) - context_length - 1)
|
||||||
|
for _ in range(batch_size)
|
||||||
|
]
|
||||||
|
input_tokens = [
|
||||||
|
encoded_text[start : start + context_length] for start in starts
|
||||||
|
]
|
||||||
|
target_tokens = [
|
||||||
|
encoded_text[start + 1 : start + context_length + 1]
|
||||||
|
for start in starts
|
||||||
|
]
|
||||||
|
return torch.tensor(input_tokens, dtype=torch.long, device=device), torch.tensor(
|
||||||
|
target_tokens, dtype=torch.long, device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def language_model_loss(logits, target_tokens):
|
||||||
|
vocabulary_size = logits.shape[-1]
|
||||||
|
return F.cross_entropy(
|
||||||
|
logits.reshape(-1, vocabulary_size), target_tokens.reshape(-1)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def generate_text(model, tokenizer, prompt, token_count):
|
||||||
|
generated_tokens = tokenizer.encode(prompt)
|
||||||
|
model.eval()
|
||||||
|
for _ in range(token_count):
|
||||||
|
context = generated_tokens[-model.config.context_length :]
|
||||||
|
input_tokens = torch.tensor([context], dtype=torch.long, device=device)
|
||||||
|
next_token_logits = model(input_tokens)[:, -1, :]
|
||||||
|
next_token = torch.categorical(next_token_logits).item()
|
||||||
|
generated_tokens.append(next_token)
|
||||||
|
return tokenizer.decode(generated_tokens)
|
||||||
|
|
||||||
|
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("steps", nargs="?", type=int, default=20)
|
||||||
|
arguments = parser.parse_args()
|
||||||
|
|
||||||
|
training_text = load_tiny_stories()
|
||||||
|
tokenizer = CharacterTokenizer(training_text)
|
||||||
|
config = TransformerConfig(vocabulary_size=len(tokenizer))
|
||||||
|
encoded_text = tokenizer.encode(training_text)
|
||||||
|
model = TinyStoriesTransformer(config).to(device)
|
||||||
|
optimizer = optim.AdamW(model.parameters(), lr=3e-4)
|
||||||
|
trainer = torch.Trainer(model, optimizer, language_model_loss, compile=True)
|
||||||
|
|
||||||
|
for step in range(arguments.steps):
|
||||||
|
input_tokens, target_tokens = create_batch(
|
||||||
|
encoded_text, batch_size=8, context_length=config.context_length
|
||||||
|
)
|
||||||
|
loss = trainer.step(input_tokens, target_tokens)
|
||||||
|
print(f"step {step + 1}: loss {loss.item():.4f}")
|
||||||
|
|
||||||
|
print(generate_text(model, tokenizer, "Once upon a time", token_count=120))
|
||||||
Reference in new issue
Block a user