| 148 | |
| 149 | |
| 150 | def train_one_epoch(config, model, criterion, data_loader, optimizer, epoch, mixup_fn, lr_scheduler): |
| 151 | model.train() |
| 152 | optimizer.zero_grad() |
| 153 | |
| 154 | num_steps = len(data_loader) |
| 155 | batch_time = AverageMeter() |
| 156 | loss_meter = AverageMeter() |
| 157 | norm_meter = AverageMeter() |
| 158 | |
| 159 | start = time.time() |
| 160 | end = time.time() |
| 161 | for idx, (samples, targets) in enumerate(data_loader): |
| 162 | samples = samples.cuda(non_blocking=True) |
| 163 | targets = targets.cuda(non_blocking=True) |
| 164 | |
| 165 | if mixup_fn is not None: |
| 166 | samples, targets = mixup_fn(samples, targets) |
| 167 | |
| 168 | outputs = model(samples) |
| 169 | |
| 170 | if config.TRAIN.ACCUMULATION_STEPS > 1: |
| 171 | loss = criterion(outputs, targets) |
| 172 | loss = loss / config.TRAIN.ACCUMULATION_STEPS |
| 173 | if config.AMP_OPT_LEVEL != "O0": |
| 174 | with amp.scale_loss(loss, optimizer) as scaled_loss: |
| 175 | scaled_loss.backward() |
| 176 | if config.TRAIN.CLIP_GRAD: |
| 177 | grad_norm = torch.nn.utils.clip_grad_norm_(amp.master_params(optimizer), config.TRAIN.CLIP_GRAD) |
| 178 | else: |
| 179 | grad_norm = get_grad_norm(amp.master_params(optimizer)) |
| 180 | else: |
| 181 | loss.backward() |
| 182 | if config.TRAIN.CLIP_GRAD: |
| 183 | grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), config.TRAIN.CLIP_GRAD) |
| 184 | else: |
| 185 | grad_norm = get_grad_norm(model.parameters()) |
| 186 | if (idx + 1) % config.TRAIN.ACCUMULATION_STEPS == 0: |
| 187 | optimizer.step() |
| 188 | optimizer.zero_grad() |
| 189 | lr_scheduler.step_update(epoch * num_steps + idx) |
| 190 | else: |
| 191 | loss = criterion(outputs, targets) |
| 192 | optimizer.zero_grad() |
| 193 | if config.AMP_OPT_LEVEL != "O0": |
| 194 | with amp.scale_loss(loss, optimizer) as scaled_loss: |
| 195 | scaled_loss.backward() |
| 196 | if config.TRAIN.CLIP_GRAD: |
| 197 | grad_norm = torch.nn.utils.clip_grad_norm_(amp.master_params(optimizer), config.TRAIN.CLIP_GRAD) |
| 198 | else: |
| 199 | grad_norm = get_grad_norm(amp.master_params(optimizer)) |
| 200 | else: |
| 201 | loss.backward() |
| 202 | if config.TRAIN.CLIP_GRAD: |
| 203 | grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), config.TRAIN.CLIP_GRAD) |
| 204 | else: |
| 205 | grad_norm = get_grad_norm(model.parameters()) |
| 206 | optimizer.step() |
| 207 | lr_scheduler.step_update(epoch * num_steps + idx) |