Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -83,8 +83,8 @@ scripts/
src/
├── data_preparation/
├── model_files/
├── pre_training/
├── models/
├── pretraining/
├── finetuning/
└── paths.py
```
Expand Down
File renamed without changes.
26 changes: 0 additions & 26 deletions configs/config.py

This file was deleted.

39 changes: 39 additions & 0 deletions configs/model.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,39 @@
from dataclasses import dataclass


@dataclass
class ModelConfig:
"""Hyperparameters for the BetterGPT model architecture.

Validated at construction time: emb_dim must be divisible by head_count,
and the resulting head_dim must be even (required by RoPE).
"""
vocab_size: int = 8192
emb_dim: int = 512
num_blocks: int = 8
head_count: int = 8
seq_length: int = 512
ffn_multiple: int = 128

def __post_init__(self):
if self.vocab_size <= 0:
raise ValueError(f"vocab_size must be > 0, got {self.vocab_size}")
if self.emb_dim <= 0:
raise ValueError(f"emb_dim must be > 0, got {self.emb_dim}")
if self.head_count <= 0:
raise ValueError(f"head_count must be > 0, got {self.head_count}")
if self.emb_dim % self.head_count != 0:
raise ValueError(
f"emb_dim ({self.emb_dim}) must be divisible by head_count ({self.head_count})"
)
head_dim = self.emb_dim // self.head_count
if head_dim % 2 != 0:
raise ValueError(
f"head_dim ({head_dim}) must be even for RoPE; adjust emb_dim or head_count"
)
if self.num_blocks <= 0:
raise ValueError(f"num_blocks must be > 0, got {self.num_blocks}")
if self.seq_length <= 0:
raise ValueError(f"seq_length must be > 0, got {self.seq_length}")
if self.ffn_multiple <= 0:
raise ValueError(f"ffn_multiple must be > 0, got {self.ffn_multiple}")
26 changes: 26 additions & 0 deletions configs/training.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
import torch
from dataclasses import dataclass

from src.paths import CHECKPOINT_DIR

checkpoint_dir = CHECKPOINT_DIR


@dataclass
class TrainingConfig:
"""Hyperparameters and runtime settings for pretraining."""
token_count: int = 470_000_000
batch_size: int = 64
learning_rate: float = 3e-4
betas: tuple = (0.9, 0.95)
eps: float = 1e-8
device: str = "cuda" if torch.cuda.is_available() else "cpu"
checkpoint_dir: str = checkpoint_dir

def __post_init__(self):
if self.token_count <= 0:
raise ValueError(f"token_count must be > 0, got {self.token_count}")
if self.batch_size < 1:
raise ValueError(f"batch_size must be >= 1, got {self.batch_size}")
if self.learning_rate <= 0:
raise ValueError(f"learning_rate must be > 0, got {self.learning_rate}")
78 changes: 0 additions & 78 deletions scripts/check_pre_trained_prediction.py

This file was deleted.

112 changes: 112 additions & 0 deletions scripts/sample_pretrain.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,112 @@
import torch
import torch.nn.functional as F
from tokenizers import Tokenizer

from configs.model import ModelConfig
from configs.training import TrainingConfig
from src.models.model import BetterGPT
from src.logger import get_logger

logger = get_logger("sample")


def load_and_predict(text, model_path, tokenizer_path, device,
max_tokens=100, temp=0.4, top_k=7, stop_at_eos=True):
"""Load a trained checkpoint and generate text for one or more prompts.

All prompts in `text` must encode to the same token length because the model
uses batched generation without a key-padding mask. For prompts of different
lengths, call this function once per prompt.

Args:
text: A single prompt string or a list of equal-length prompt strings.
model_path: Path to the .pt checkpoint file produced by the trainer.
tokenizer_path: Path to the tokenizer JSON file.
device: Torch device string ('cpu', 'cuda', etc.).
max_tokens: Maximum number of new tokens to generate per prompt.
temp: Sampling temperature; 0 for greedy decoding.
top_k: Restrict sampling to the top-k logits; None for unrestricted.
stop_at_eos: Stop generation when the <eos> token is produced.

Returns:
List of decoded output strings, one per prompt.
"""
if max_tokens <= 0:
raise ValueError(f"max_tokens must be > 0, got {max_tokens}")
if temp < 0:
raise ValueError(f"temp must be >= 0, got {temp}")
if top_k is not None and top_k <= 0:
raise ValueError(f"top_k must be > 0, got {top_k}")

tokenizer = Tokenizer.from_file(tokenizer_path)
eos_id = tokenizer.token_to_id("<eos>") if stop_at_eos else None

ckpt = torch.load(model_path, map_location=device, weights_only=True)
model = BetterGPT(ModelConfig(**ckpt["model_config"]))
model.load_state_dict(ckpt["model_state"])
model = model.to(device)
logger.info(
"Model loaded from %s | vocab_size=%d | emb_dim=%d | num_blocks=%d",
model_path,
ckpt["model_config"].get("vocab_size"),
ckpt["model_config"].get("emb_dim"),
ckpt["model_config"].get("num_blocks"),
)

if isinstance(text, str):
text = [text]
enc_batch = tokenizer.encode_batch(text)

seqs = []
for enc in enc_batch:
ids = enc.ids
if eos_id is not None and ids and ids[-1] == eos_id:
ids = ids[:-1] # drop trailing eos token
seqs.append(ids)

lengths = {len(s) for s in seqs}
assert len(lengths) == 1, (
f"prompts differ in length {lengths}; batched gen needs a key-padding "
f"mask the model lacks. Pass equal-length prompts, or generate one at a time."
)

idx = torch.tensor(seqs, device=device)
out = model.generate(idx, max_tokens=max_tokens, temp=temp,
top_k=top_k, eos_id=eos_id)
return tokenizer.decode_batch(out.tolist(), skip_special_tokens=True)


if __name__ == "__main__":
import argparse
import os

training_config = TrainingConfig()
path = training_config.checkpoint_dir
model_path = os.path.join(path, "best_model.pt")

if not os.path.exists(model_path):
raise FileNotFoundError(f"Model checkpoint not found at {model_path}")

tokenizer_path = "./tokenizer.json"
device = training_config.device

parser = argparse.ArgumentParser()
parser.add_argument("--text", type=str, required=True)
parser.add_argument("--max_tokens", type=int, default=500)
parser.add_argument("--temp", type=float, default=0.4)
parser.add_argument("--top_k", type=int, default=7)
parser.add_argument("--stop_at_eos", action="store_true", default=True)

args = parser.parse_args()

output = load_and_predict(
text=args.text,
model_path=model_path,
tokenizer_path=tokenizer_path,
device=device,
max_tokens=args.max_tokens,
temp=args.temp,
top_k=args.top_k,
stop_at_eos=args.stop_at_eos,
)
print(output)
Loading
Loading