()
| 249 | |
| 250 | |
| 251 | def train(): |
| 252 | global local_rank |
| 253 | |
| 254 | parser = transformers.HfArgumentParser( |
| 255 | (ModelArguments, DataArguments, TrainingArguments) |
| 256 | ) |
| 257 | model_args, data_args, training_args = parser.parse_args_into_dataclasses() |
| 258 | if training_args.source_model_max_length < training_args.model_max_length: |
| 259 | condense_ratio = int(training_args.model_max_length/training_args.source_model_max_length) |
| 260 | # ratio = N means the sequence length is expanded by N, remember to change the model_max_length to 8192 (2048 * ratio) for ratio = 4 |
| 261 | replace_llama_with_condense(ratio=condense_ratio) |
| 262 | local_rank = training_args.local_rank |
| 263 | tokenizer = transformers.AutoTokenizer.from_pretrained( |
| 264 | model_args.model_name_or_path, |
| 265 | cache_dir=training_args.cache_dir, |
| 266 | model_max_length=training_args.model_max_length, |
| 267 | padding_side="right", |
| 268 | use_fast=False, |
| 269 | ) |
| 270 | tokenizer.pad_token = tokenizer.unk_token |
| 271 | |
| 272 | data_module = make_supervised_data_module(tokenizer=tokenizer, data_args=data_args) |
| 273 | world_size = int(os.environ.get("WORLD_SIZE", 1)) |
| 274 | ddp = world_size != 1 |
| 275 | device_map = {"": int(os.environ.get("LOCAL_RANK") or 0)} if ddp else None |
| 276 | model = transformers.AutoModelForCausalLM.from_pretrained( |
| 277 | model_args.model_name_or_path, |
| 278 | cache_dir=training_args.cache_dir, |
| 279 | device_map=device_map |
| 280 | ) |
| 281 | model.config.use_cache = False |
| 282 | trainer = Trainer( |
| 283 | model=model, tokenizer=tokenizer, args=training_args, **data_module |
| 284 | ) |
| 285 | |
| 286 | if list(pathlib.Path(training_args.output_dir).glob("checkpoint-*")): |
| 287 | trainer.train(resume_from_checkpoint=True) |
| 288 | else: |
| 289 | trainer.train() |
| 290 | trainer.save_state() |
| 291 | safe_save_model_for_hf_trainer(trainer=trainer, output_dir=training_args.output_dir) |
| 292 | |
| 293 | |
| 294 | if __name__ == "__main__": |
no test coverage detected