File size: 1,511 Bytes
701cf7d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 | """
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()
|