Ares_v1 / scripts /pretrain.py
Ares Deployer
Deploy Ares full from scratch: BPE 128K, RoPE 8192, GQA+KV, RMSNorm, SwiGLU, RAG SQLite, CoT/ToT/Planner, SFT/RLHF, code+search
701cf7d
Raw
History Blame Contribute Delete
2.55 kB
"""
Pretrain Ares from scratch - iterative scaling
"""
import sys
sys.path.append("src")
import os
import torch
from ares.config import get_config
from ares.model.model import AresForCausalLM
from ares.tokenizer.tokenizer import AresTokenizer
from ares.training.dataset_pipeline import DataPipeline
from ares.training.trainer import AresTrainer
from torch.utils.data import Dataset, DataLoader
class TokenizedDataset(Dataset):
def __init__(self, token_chunks):
self.chunks = token_chunks
def __len__(self): return len(self.chunks)
def __getitem__(self, idx):
chunk = self.chunks[idx]
# input = chunk, label = chunk (shift inside model)
return torch.tensor(chunk, dtype=torch.long), torch.tensor(chunk, dtype=torch.long)
def main():
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--config", type=str, default="tiny", choices=["tiny","small","medium","billion","large"])
parser.add_argument("--tokenizer", type=str, default="data/tokenizer.json")
parser.add_argument("--steps", type=int, default=500)
parser.add_argument("--max_seq", type=int, default=512)
args = parser.parse_args()
config = get_config(args.config)
print(f"[Ares] Config {args.config}: ~{config.num_parameters_approx/1e6:.1f}M params")
tokenizer = AresTokenizer(vocab_file=args.tokenizer if os.path.exists(args.tokenizer) else None, vocab_size=config.vocab_size)
print(f"[Ares] Tokenizer vocab {len(tokenizer.vocab)}")
model = AresForCausalLM(config)
print(f"[Ares] Actual params {model.count_parameters()/1e6:.2f}M")
pipeline = DataPipeline(tokenizer, max_seq_len=args.max_seq)
texts = pipeline.load_hf_dataset("allenai/c4", num_samples=20000)
# Limit chunks for demo
chunks = []
for chunk in pipeline.tokenized_stream(texts):
chunks.append(chunk)
if len(chunks) >= 5000:
break
print(f"[Data] Got {len(chunks)} chunks")
dataset = TokenizedDataset(chunks)
# DataLoader that yields tensors
def collate(batch):
input_ids = torch.stack([b[0] for b in batch])
labels = torch.stack([b[1] for b in batch])
return input_ids, labels
loader = DataLoader(dataset, batch_size=2, shuffle=True, collate_fn=collate)
trainer = AresTrainer(model, tokenizer, config, device="cuda" if torch.cuda.is_available() else "cpu")
trainer.train(loader, total_steps=args.steps, warmup_steps=50, lr=3e-4, save_dir=f"checkpoints/{args.config}")
if __name__ == "__main__":
main()