()
| 100 | return dict(train_dataset=train_dataset, data_collator=data_collator) |
| 101 | |
| 102 | def train(): |
| 103 | parser = transformers.HfArgumentParser((ModelArguments, DataArguments, TrainingArguments)) |
| 104 | model_args, data_args, training_args = parser.parse_args_into_dataclasses() |
| 105 | if "chatglm" in model_args.model_name_or_path.lower() or "longalign-6b" in model_args.model_name_or_path.lower(): |
| 106 | model = AutoModelForCausalLM.from_pretrained( |
| 107 | model_args.model_name_or_path, |
| 108 | torch_dtype=torch.bfloat16, |
| 109 | trust_remote_code=True, empty_init=False |
| 110 | ) |
| 111 | tokenizer = AutoTokenizer.from_pretrained( |
| 112 | model_args.model_name_or_path, |
| 113 | trust_remote_code=True |
| 114 | ) |
| 115 | else: |
| 116 | model = AutoModelForCausalLM.from_pretrained(model_args.model_name_or_path, |
| 117 | torch_dtype=torch.bfloat16, |
| 118 | trust_remote_code=True) |
| 119 | tokenizer = AutoTokenizer.from_pretrained(model_args.model_name_or_path, |
| 120 | trust_remote_code=True) |
| 121 | |
| 122 | if training_args.lora_enable == True: |
| 123 | lora_config = LoraConfig( |
| 124 | r=training_args.lora_rank, |
| 125 | lora_alpha=training_args.lora_alpha, |
| 126 | target_modules='all-linear', |
| 127 | lora_dropout=training_args.lora_dropout, |
| 128 | bias='none', |
| 129 | task_type="CAUSAL_LM", |
| 130 | ) |
| 131 | if training_args.bf16: |
| 132 | model.to(torch.bfloat16) |
| 133 | if training_args.fp16: |
| 134 | model.to(torch.float16) |
| 135 | #gradient checkpointing |
| 136 | if training_args.gradient_checkpointing: |
| 137 | model.enable_input_require_grads() |
| 138 | model = get_peft_model(model, lora_config) |
| 139 | |
| 140 | if model_args.pack_loss: |
| 141 | model.pack_loss = True |
| 142 | data_module = make_supervised_data_module(data_args=data_args) |
| 143 | |
| 144 | trainer = TrainerNoShuffle( |
| 145 | model=model, |
| 146 | tokenizer=tokenizer, |
| 147 | args=training_args, |
| 148 | **data_module |
| 149 | ) |
| 150 | |
| 151 | trainer.train(resume_from_checkpoint=False) |
| 152 | trainer.save_model() |
| 153 | |
| 154 | if __name__ == "__main__": |
| 155 | train() |
no test coverage detected