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