Text Generation
Transformers
finance
domain-specialization
amd-rocm
mi300x
fine-tuned
ml-intern
amd-finance-llm / train.py
shah-shazid-askary's picture
Upload train.py with huggingface_hub
8fd20e0 verified
Raw History Blame Contribute Delete
6.04 kB
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()