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