Ares_v1 / scripts /sft_train.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
1.51 kB
"""
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()