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 | """ | |
| Supervised Fine-Tuning script | |
| """ | |
| import sys | |
| sys.path.append("src") | |
| 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.alignment.sft import SFTDataset, SFTTrainer | |
| import os | |
| def main(): | |
| import argparse | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--config", type=str, default="tiny") | |
| parser.add_argument("--tokenizer", type=str, default="data/tokenizer.json") | |
| parser.add_argument("--model_path", type=str, default=None) | |
| parser.add_argument("--epochs", type=int, default=1) | |
| args = parser.parse_args() | |
| config = get_config(args.config) | |
| tokenizer = AresTokenizer(vocab_file=args.tokenizer if os.path.exists(args.tokenizer) else None, vocab_size=config.vocab_size) | |
| model = AresForCausalLM(config) | |
| if args.model_path and os.path.exists(args.model_path): | |
| model.load_state_dict(torch.load(args.model_path, map_location="cpu")) | |
| print(f"Loaded {args.model_path}") | |
| pipeline = DataPipeline(tokenizer, max_seq_len=config.max_position_embeddings) | |
| examples = pipeline.build_sft_dataset(num_samples=5000) | |
| print(f"SFT examples {len(examples)}") | |
| ds = SFTDataset(examples, tokenizer, max_len=1024) | |
| trainer = SFTTrainer(model, tokenizer) | |
| trainer.train(ds, epochs=args.epochs, lr=2e-5, batch_size=2, save_dir="checkpoints_sft") | |
| if __name__ == "__main__": | |
| import torch | |
| main() | |