humanist96/FinGPT_Forecaster
0
1from transformers.integrations import TensorBoardCallback2from transformers import AutoTokenizer, AutoModel, AutoModelForCausalLM3from transformers import TrainingArguments, Trainer, DataCollatorForSeq2Seq4from transformers import TrainerCallback, TrainerState, TrainerControl5from transformers.trainer import TRAINING_ARGS_NAME6from torch.utils.tensorboard import SummaryWriter7import datasets8import torch9import os10import re11import sys12import wandb13import argparse14from datetime import datetime15from functools import partial16from tqdm import tqdm17from utils import *18 19# LoRA20from peft import (21 TaskType,22 LoraConfig,23 get_peft_model,24 get_peft_model_state_dict,25 prepare_model_for_int8_training,26 set_peft_model_state_dict, 27)28 29# Replace with your own api_key and project name30os.environ['WANDB_API_KEY'] = 'ecf1e5e4f47441d46822d38a3249d62e8fc94db4'31os.environ['WANDB_PROJECT'] = 'fingpt-forecaster'32 33 34class GenerationEvalCallback(TrainerCallback):35 36 def __init__(self, eval_dataset, ignore_until_epoch=0):37 self.eval_dataset = eval_dataset38 self.ignore_until_epoch = ignore_until_epoch39 40 def on_evaluate(self, args, state: TrainerState, control: TrainerControl, **kwargs):41 42 if state.epoch is None or state.epoch + 1 < self.ignore_until_epoch:43 return44 45 if state.is_local_process_zero:46 model = kwargs['model']47 tokenizer = kwargs['tokenizer']48 generated_texts, reference_texts = [], []49 50 for feature in tqdm(self.eval_dataset):51 prompt = feature['prompt']52 gt = feature['answer']53 inputs = tokenizer(54 prompt, return_tensors='pt',55 padding=False, max_length=409656 )57 inputs = {key: value.to(model.device) for key, value in inputs.items()}58 59 res = model.generate(60 **inputs, 61 use_cache=True62 )63 output = tokenizer.decode(res[0], skip_special_tokens=True)64 answer = re.sub(r'.*\[/INST\]\s*', '', output, flags=re.DOTALL)65 66 generated_texts.append(answer)67 reference_texts.append(gt)68 69 # print("GENERATED: ", answer)70 # print("REFERENCE: ", gt)71 72 metrics = calc_metrics(reference_texts, generated_texts)73 74 # Ensure wandb is initialized75 if wandb.run is None:76 wandb.init()77 78 wandb.log(metrics, step=state.global_step)79 torch.cuda.empty_cache() 80 81 82def main(args):83 84 model_name = parse_model_name(args.base_model, args.from_remote)85 86 # load model87 model = AutoModelForCausalLM.from_pretrained(88 model_name,89 # load_in_8bit=True,90 trust_remote_code=True91 )92 if args.local_rank == 0:93 print(model)94 95 tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)96 tokenizer.pad_token = tokenizer.eos_token97 tokenizer.padding_side = "right"98 99 # load data100 dataset_list = load_dataset(args.dataset, args.from_remote)101 102 dataset_train = datasets.concatenate_datasets([d['train'] for d in dataset_list]).shuffle(seed=42)103 104 if args.test_dataset:105 dataset_list = load_dataset(args.test_dataset, args.from_remote)106 107 dataset_test = datasets.concatenate_datasets([d['test'] for d in dataset_list])108 109 original_dataset = datasets.DatasetDict({'train': dataset_train, 'test': dataset_test})110 111 eval_dataset = original_dataset['test'].shuffle(seed=42).select(range(50))112 113 dataset = original_dataset.map(partial(tokenize, args, tokenizer))114 print('original dataset length: ', len(dataset['train']))115 dataset = dataset.filter(lambda x: not x['exceed_max_length'])116 print('filtered dataset length: ', len(dataset['train']))117 dataset = dataset.remove_columns(118 ['prompt', 'answer', 'label', 'symbol', 'period', 'exceed_max_length']119 )120 121 current_time = datetime.now()122 formatted_time = current_time.strftime('%Y%m%d%H%M')123 124 training_args = TrainingArguments(125 output_dir=f'finetuned_models/{args.run_name}_{formatted_time}', # 保存位置126 logging_steps=args.log_interval,127 num_train_epochs=args.num_epochs,128 per_device_train_batch_size=args.batch_size,129 per_device_eval_batch_size=args.batch_size,130 gradient_accumulation_steps=args.gradient_accumulation_steps,131 dataloader_num_workers=args.num_workers,132 learning_rate=args.learning_rate,133 weight_decay=args.weight_decay,134 warmup_ratio=args.warmup_ratio,135 lr_scheduler_type=args.scheduler,136 save_steps=args.eval_steps,137 eval_steps=args.eval_steps,138 fp16=True,139 deepspeed=args.ds_config,140 evaluation_strategy=args.evaluation_strategy,141 remove_unused_columns=False,142 report_to='wandb',143 run_name=args.run_name144 )145 146 model.gradient_checkpointing_enable()147 model.enable_input_require_grads()148 model.is_parallelizable = True149 model.model_parallel = True150 model.model.config.use_cache = False151 152 # model = prepare_model_for_int8_training(model)153 154 # setup peft155 peft_config = LoraConfig(156 task_type=TaskType.CAUSAL_LM,157 inference_mode=False,158 r=8,159 lora_alpha=16,160 lora_dropout=0.1,161 target_modules=lora_module_dict[args.base_model],162 bias='none',163 )164 model = get_peft_model(model, peft_config)165 166 # Train167 trainer = Trainer(168 model=model, 169 args=training_args, 170 train_dataset=dataset['train'],171 eval_dataset=dataset['test'], 172 tokenizer=tokenizer,173 data_collator=DataCollatorForSeq2Seq(174 tokenizer, padding=True,175 return_tensors="pt"176 ),177 callbacks=[178 GenerationEvalCallback(179 eval_dataset=eval_dataset,180 ignore_until_epoch=round(0.3 * args.num_epochs)181 )182 ]183 )184 185 if torch.__version__ >= "2" and sys.platform != "win32":186 model = torch.compile(model)187 188 torch.cuda.empty_cache()189 trainer.train()190 191 # save model192 model.save_pretrained(training_args.output_dir)193 194 195if __name__ == "__main__":196 197 parser = argparse.ArgumentParser()198 parser.add_argument("--local_rank", default=0, type=int)199 parser.add_argument("--run_name", default='local-test', type=str)200 parser.add_argument("--dataset", required=True, type=str)201 parser.add_argument("--test_dataset", type=str)202 parser.add_argument("--base_model", required=True, type=str, choices=['chatglm2', 'llama2'])203 parser.add_argument("--max_length", default=512, type=int)204 parser.add_argument("--batch_size", default=4, type=int, help="The train batch size per device")205 parser.add_argument("--learning_rate", default=1e-4, type=float, help="The learning rate")206 parser.add_argument("--weight_decay", default=0.01, type=float, help="weight decay")207 parser.add_argument("--num_epochs", default=8, type=float, help="The training epochs")208 parser.add_argument("--num_workers", default=8, type=int, help="dataloader workers")209 parser.add_argument("--log_interval", default=20, type=int)210 parser.add_argument("--gradient_accumulation_steps", default=8, type=int)211 parser.add_argument("--warmup_ratio", default=0.05, type=float)212 parser.add_argument("--ds_config", default='./config_new.json', type=str)213 parser.add_argument("--scheduler", default='linear', type=str)214 parser.add_argument("--instruct_template", default='default')215 parser.add_argument("--evaluation_strategy", default='steps', type=str) 216 parser.add_argument("--eval_steps", default=0.1, type=float) 217 parser.add_argument("--from_remote", default=False, type=bool) 218 args = parser.parse_args()219 220 wandb.login()221 main(args)