mirror of
https://github.com/priyanshujain/torchmlx.git
synced 2026-10-02 11:07:13 +00:00
match notebook training config
This commit is contained in:
1 parent
0c1e9fcb84
commit
cf90000071
1 file changed
+26
-14
@@ -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))
|
||||||
Reference in new issue
Block a user