Train the model
(args, train_dataset, model, tokenizer)
| 77 | |
| 78 | |
| 79 | def train(args, train_dataset, model, tokenizer): |
| 80 | """ Train the model """ |
| 81 | if args.local_rank in [-1, 0]: |
| 82 | tb_writer = SummaryWriter() |
| 83 | |
| 84 | args.train_batch_size = args.per_gpu_train_batch_size * max(1, args.n_gpu) |
| 85 | train_sampler = RandomSampler(train_dataset) if args.local_rank == -1 else DistributedSampler(train_dataset) |
| 86 | train_dataloader = DataLoader(train_dataset, sampler=train_sampler, batch_size=args.train_batch_size) |
| 87 | |
| 88 | if args.max_steps > 0: |
| 89 | t_total = args.max_steps |
| 90 | args.num_train_epochs = args.max_steps // (len(train_dataloader) // args.gradient_accumulation_steps) + 1 |
| 91 | else: |
| 92 | t_total = len(train_dataloader) // args.gradient_accumulation_steps * args.num_train_epochs |
| 93 | |
| 94 | # Prepare optimizer and schedule (linear warmup and decay) |
| 95 | no_decay = ["bias", "LayerNorm.weight"] |
| 96 | optimizer_grouped_parameters = [ |
| 97 | { |
| 98 | "params": [p for n, p in model.named_parameters() if not any(nd in n for nd in no_decay)], |
| 99 | "weight_decay": args.weight_decay, |
| 100 | }, |
| 101 | {"params": [p for n, p in model.named_parameters() if any(nd in n for nd in no_decay)], "weight_decay": 0.0}, |
| 102 | ] |
| 103 | optimizer = AdamW(optimizer_grouped_parameters, lr=args.learning_rate, eps=args.adam_epsilon) |
| 104 | scheduler = get_linear_schedule_with_warmup( |
| 105 | optimizer, num_warmup_steps=args.warmup_steps, num_training_steps=t_total |
| 106 | ) |
| 107 | if args.fp16: |
| 108 | try: |
| 109 | from apex import amp |
| 110 | except ImportError: |
| 111 | raise ImportError("Please install apex from https://www.github.com/nvidia/apex to use fp16 training.") |
| 112 | model, optimizer = amp.initialize(model, optimizer, opt_level=args.fp16_opt_level) |
| 113 | |
| 114 | # multi-gpu training (should be after apex fp16 initialization) |
| 115 | if args.n_gpu > 1: |
| 116 | model = torch.nn.DataParallel(model) |
| 117 | |
| 118 | # Distributed training (should be after apex fp16 initialization) |
| 119 | if args.local_rank != -1: |
| 120 | model = torch.nn.parallel.DistributedDataParallel( |
| 121 | model, device_ids=[args.local_rank], output_device=args.local_rank, find_unused_parameters=True |
| 122 | ) |
| 123 | |
| 124 | # Train! |
| 125 | logger.info("***** Running training *****") |
| 126 | logger.info(" Num examples = %d", len(train_dataset)) |
| 127 | logger.info(" Num Epochs = %d", args.num_train_epochs) |
| 128 | logger.info(" Instantaneous batch size per GPU = %d", args.per_gpu_train_batch_size) |
| 129 | logger.info( |
| 130 | " Total train batch size (w. parallel, distributed & accumulation) = %d", |
| 131 | args.train_batch_size |
| 132 | * args.gradient_accumulation_steps |
| 133 | * (torch.distributed.get_world_size() if args.local_rank != -1 else 1), |
| 134 | ) |
| 135 | logger.info(" Gradient Accumulation steps = %d", args.gradient_accumulation_steps) |
| 136 | logger.info(" Total optimization steps = %d", t_total) |
no test coverage detected