MCPcopy Create free account
hub / github.com/SwinTransformer/Transformer-SSL / train_one_epoch

Function train_one_epoch

main.py:150–228  ·  view source on GitHub ↗
(config, model, criterion, data_loader, optimizer, epoch, mixup_fn, lr_scheduler)

Source from the content-addressed store, hash-verified

148
149
150def 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)

Callers 1

mainFunction · 0.70

Calls 1

get_grad_normFunction · 0.90

Tested by

no test coverage detected