""" 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()