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