Text Generation
Transformers
finance
domain-specialization
amd-rocm
mi300x
fine-tuned
ml-intern
File size: 6,041 Bytes
8fd20e0
dd7a648
8fd20e0
dd7a648
 
 
 
 
8fd20e0
dd7a648
8fd20e0
dd7a648
8fd20e0
dd7a648
 
 
8fd20e0
dd7a648
8fd20e0
 
 
 
dd7a648
 
 
8fd20e0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dd7a648
 
8fd20e0
dd7a648
 
8fd20e0
dd7a648
 
 
 
8fd20e0
 
dd7a648
 
 
 
8fd20e0
dd7a648
8fd20e0
 
dd7a648
8fd20e0
 
 
dd7a648
8fd20e0
 
dd7a648
 
 
 
8fd20e0
dd7a648
8fd20e0
 
 
dd7a648
 
 
 
8fd20e0
dd7a648
 
8fd20e0
 
 
 
 
 
 
 
 
 
 
 
dd7a648
 
 
8fd20e0
 
 
dd7a648
8fd20e0
dd7a648
8fd20e0
dd7a648
 
 
 
8fd20e0
dd7a648
8fd20e0
dd7a648
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
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
import os, sys, argparse, logging
import torch
from datasets import load_dataset
from transformers import AutoModelForCausalLM, AutoTokenizer
from trl import SFTTrainer, SFTConfig
from peft import LoraConfig, TaskType

if torch.cuda.is_available():
    print(f'ROCm/GPU: {torch.cuda.is_available()}, devices={torch.cuda.device_count()}')
    for i in range(torch.cuda.device_count()):
        print(f'  Device {i}: {torch.cuda.get_device_name(i)}')
else:
    print('WARNING: No GPU — training will be extremely slow.')

def to_messages(example):
    messages = []
    system = example.get('system', '')
    if system and str(system).strip():
        messages.append({'role': 'system', 'content': str(system).strip()})
    messages.append({'role': 'user', 'content': str(example['user']).strip()})
    messages.append({'role': 'assistant', 'content': str(example['assistant']).strip()})
    return {'messages': messages}

def main():
    parser = argparse.ArgumentParser()
    parser.add_argument('--model_name', default='Qwen/Qwen2.5-7B-Instruct')
    parser.add_argument('--dataset_name', default='Josephgflowers/Finance-Instruct-500k')
    parser.add_argument('--output_dir', default='/app/output/finance-sft')
    parser.add_argument('--hub_model_id', default='shah-shazid-askary/amd-finance-llm')
    parser.add_argument('--learning_rate', type=float, default=1e-5)
    parser.add_argument('--num_train_epochs', type=int, default=3)
    parser.add_argument('--warmup_ratio', type=float, default=0.1)
    parser.add_argument('--max_seq_length', type=int, default=8192)
    parser.add_argument('--per_device_train_batch_size', type=int, default=1)
    parser.add_argument('--gradient_accumulation_steps', type=int, default=16)
    parser.add_argument('--bf16', action='store_true', default=True)
    parser.add_argument('--fp16', action='store_true', default=False)
    parser.add_argument('--use_lora', action='store_true', default=False)
    parser.add_argument('--lora_r', type=int, default=32)
    parser.add_argument('--lora_alpha', type=int, default=16)
    parser.add_argument('--lora_dropout', type=float, default=0.05)
    parser.add_argument('--max_samples', type=int, default=None)
    parser.add_argument('--logging_steps', type=int, default=10)
    parser.add_argument('--save_steps', type=int, default=500)
    parser.add_argument('--eval_steps', type=int, default=500)
    parser.add_argument('--seed', type=int, default=42)
    parser.add_argument('--push_to_hub', action='store_true', default=True)
    parser.add_argument('--resume_from_checkpoint', type=str, default=None)
    parser.add_argument('--trackio_space_id', type=str, default=None)
    parser.add_argument('--trackio_project', type=str, default='amd-finance-llm')
    args = parser.parse_args()

    logging.basicConfig(level=logging.INFO, format='%(asctime)s %(levelname)s %(message)s')
    logger = logging.getLogger(__name__)

    logger.info('Loading tokenizer...')
    tokenizer = AutoTokenizer.from_pretrained(args.model_name, trust_remote_code=True)
    if tokenizer.pad_token is None:
        tokenizer.pad_token = tokenizer.eos_token

    logger.info('Loading dataset...')
    ds = load_dataset(args.dataset_name, split='train')
    if args.max_samples:
        ds = ds.select(range(min(args.max_samples, len(ds))))
    ds = ds.map(to_messages, remove_columns=ds.column_names, batched=False)
    ds = ds.train_test_split(test_size=0.05, seed=args.seed)
    logger.info(f'Train: {len(ds["train"])} | Eval: {len(ds["test"])}')

    logger.info('Loading model...')
    attn_impl = 'flash_attention_2' if torch.cuda.is_available() else 'eager'
    model_kwargs = {
        'torch_dtype': torch.bfloat16 if args.bf16 else (torch.float16 if args.fp16 else torch.float32),
        'attn_implementation': attn_impl,
        'trust_remote_code': True,
    }
    if torch.cuda.device_count() >= 1 and not os.environ.get('ACCELERATE_USE_DEEPSPEED'):
        model_kwargs['device_map'] = 'auto'
    model = AutoModelForCausalLM.from_pretrained(args.model_name, **model_kwargs)

    peft_config = None
    if args.use_lora:
        logger.info('Applying LoRA...')
        peft_config = LoraConfig(
            task_type=TaskType.CAUSAL_LM, inference_mode=False,
            r=args.lora_r, lora_alpha=args.lora_alpha, lora_dropout=args.lora_dropout,
            target_modules=['q_proj','k_proj','v_proj','o_proj','gate_proj','up_proj','down_proj'],
        )
        model.print_trainable_parameters()

    sft_args = SFTConfig(
        output_dir=args.output_dir, num_train_epochs=args.num_train_epochs,
        per_device_train_batch_size=args.per_device_train_batch_size,
        gradient_accumulation_steps=args.gradient_accumulation_steps,
        learning_rate=args.learning_rate, warmup_ratio=args.warmup_ratio,
        lr_scheduler_type='cosine', max_seq_length=args.max_seq_length,
        bf16=args.bf16, fp16=args.fp16,
        logging_steps=args.logging_steps, logging_strategy='steps',
        logging_first_step=True, save_steps=args.save_steps,
        eval_strategy='steps', eval_steps=args.eval_steps,
        load_best_model_at_end=True, metric_for_best_model='eval_loss',
        greater_is_better=False, seed=args.seed,
        push_to_hub=args.push_to_hub, hub_model_id=args.hub_model_id,
        disable_tqdm=True, gradient_checkpointing=True,
        report_to=['trackio'] if args.trackio_space_id else None,
        run_name=f'finance-sft-{args.model_name.split("/")[-1]}-lr{args.learning_rate}',
    )

    trainer = SFTTrainer(
        model=model, tokenizer=tokenizer,
        train_dataset=ds['train'], eval_dataset=ds['test'],
        args=sft_args, peft_config=peft_config,
    )
    logger.info('Starting training...')
    trainer.train(resume_from_checkpoint=args.resume_from_checkpoint)
    logger.info('Saving model...')
    trainer.save_model(args.output_dir)
    tokenizer.save_pretrained(args.output_dir)
    if args.push_to_hub:
        trainer.push_to_hub()
    logger.info('Done.')

if __name__ == '__main__':
    main()