match notebook training config

This commit is contained in:
pj committed 2026-09-19 23:25:48 +05:30
1 parent 0c1e9fcb84
commit cf90000071
1 file changed
+26 -14
+26 -14
View File
@@ -1,4 +1,3 @@
import argparse
import random import random
import urllib.request import urllib.request
from dataclasses import dataclass from dataclasses import dataclass
@@ -24,11 +23,20 @@ DATA_PATH = Path(__file__).with_name("TinyStoriesV2-GPT4-valid.txt")
@dataclass @dataclass
class TransformerConfig: class TransformerConfig:
vocabulary_size: int vocabulary_size: int
context_length: int = 64 context_length: int = 128
model_dimension: int = 64 model_dimension: int = 32
head_count: int = 4 head_count: int = 16
layer_count: int = 2 layer_count: int = 4
feed_forward_dimension: int = 256 feed_forward_dimension: int = 128
@dataclass
class TrainingConfig:
batch_size: int = 32
step_count: int = 5000
learning_rate: float = 1e-3
weight_decay: float = 1e-2
log_interval: int = 100
class CharacterTokenizer: class CharacterTokenizer:
@@ -225,27 +233,31 @@ def generate_text(model, tokenizer, prompt, token_count):
return tokenizer.decode(generated_tokens) 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() training_text = load_tiny_stories()
tokenizer = CharacterTokenizer(training_text) tokenizer = CharacterTokenizer(training_text)
config = TransformerConfig(vocabulary_size=len(tokenizer)) config = TransformerConfig(vocabulary_size=len(tokenizer))
training_config = TrainingConfig()
encoded_text = tokenizer.encode(training_text) encoded_text = tokenizer.encode(training_text)
model = TinyStoriesTransformer(config).to(device) model = TinyStoriesTransformer(config).to(device)
optimizer = optim.AdamW(model.parameters(), lr=3e-4) optimizer = optim.AdamW(
model.parameters(),
lr=training_config.learning_rate,
weight_decay=training_config.weight_decay,
)
model.train() model.train()
for step in range(arguments.steps): for step in range(training_config.step_count):
input_tokens, target_tokens = create_batch( input_tokens, target_tokens = create_batch(
encoded_text, batch_size=8, context_length=config.context_length encoded_text,
batch_size=training_config.batch_size,
context_length=config.context_length,
) )
optimizer.zero_grad() optimizer.zero_grad()
logits = model(input_tokens) logits = model(input_tokens)
loss = language_model_loss(logits, target_tokens) loss = language_model_loss(logits, target_tokens)
loss.backward() loss.backward()
optimizer.step() optimizer.step()
print(f"step {step + 1}: loss {loss.item():.4f}") if step % training_config.log_interval == 0:
print(f"step {step}: loss {loss.item():.4f}")
print(generate_text(model, tokenizer, "Once upon a time", token_count=120)) print(generate_text(model, tokenizer, "Once upon a time", token_count=120))