()
| 89 | |
| 90 | |
| 91 | def train(): |
| 92 | parser = transformers.HfArgumentParser( |
| 93 | (ModelArguments, DataArguments, TrainingArguments, LoraArguments) |
| 94 | ) |
| 95 | ( |
| 96 | model_args, |
| 97 | data_args, |
| 98 | training_args, |
| 99 | lora_args, |
| 100 | ) = parser.parse_args_into_dataclasses() |
| 101 | |
| 102 | if training_args.source_model_max_length < training_args.model_max_length: |
| 103 | condense_ratio = int(training_args.model_max_length/training_args.source_model_max_length) |
| 104 | # ratio = N means the sequence length is expanded by N, remember to change the model_max_length to 8192 (2048 * ratio) for ratio = 4 |
| 105 | replace_llama_with_condense(ratio=condense_ratio) |
| 106 | |
| 107 | world_size = int(os.environ.get("WORLD_SIZE", 1)) |
| 108 | ddp = world_size != 1 |
| 109 | device_map = {"": int(os.environ.get("LOCAL_RANK") or 0)} if ddp else None |
| 110 | model = transformers.AutoModelForCausalLM.from_pretrained( |
| 111 | model_args.model_name_or_path, |
| 112 | cache_dir=training_args.cache_dir, |
| 113 | device_map=device_map |
| 114 | ) |
| 115 | lora_config = LoraConfig( |
| 116 | r=lora_args.lora_r, |
| 117 | lora_alpha=lora_args.lora_alpha, |
| 118 | target_modules=lora_args.lora_target_modules, |
| 119 | lora_dropout=lora_args.lora_dropout, |
| 120 | bias=lora_args.lora_bias, |
| 121 | task_type="CAUSAL_LM", |
| 122 | ) |
| 123 | model = get_peft_model(model, lora_config) |
| 124 | if training_args.deepspeed is not None and training_args.local_rank == 0: |
| 125 | model.print_trainable_parameters() |
| 126 | |
| 127 | if training_args.gradient_checkpointing: |
| 128 | logging.warning( |
| 129 | "gradient checkpointing with lora makes requires_grad " |
| 130 | "incorrect and needs a monkey patch in Trainer or the " |
| 131 | "wrapped model's forward. ref: " |
| 132 | "https://github.com/lm-sys/FastChat/pull/138#issuecomment-1509172198" |
| 133 | ) |
| 134 | model.enable_input_require_grads() |
| 135 | |
| 136 | tokenizer = transformers.AutoTokenizer.from_pretrained( |
| 137 | model_args.model_name_or_path, |
| 138 | cache_dir=training_args.cache_dir, |
| 139 | model_max_length=training_args.model_max_length, |
| 140 | padding_side="right", |
| 141 | use_fast=False, |
| 142 | ) |
| 143 | tokenizer.pad_token = tokenizer.unk_token |
| 144 | |
| 145 | data_module = make_supervised_data_module(tokenizer=tokenizer, data_args=data_args) |
| 146 | trainer = Trainer( |
| 147 | model=model, tokenizer=tokenizer, args=training_args, **data_module |
| 148 | ) |
no test coverage detected