CoolFace
Apppublic

humanist96/FinGPT_Forecaster

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
train_lora.py221 linesDownload Raw Back to root
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)