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