gradient accumulation
This commit is contained in:
@@ -1,5 +1,5 @@
|
|||||||
data:
|
data:
|
||||||
batch_size: 512
|
batch_size: 16
|
||||||
color_jitter_strength: 0.0
|
color_jitter_strength: 0.0
|
||||||
crop_scale:
|
crop_scale:
|
||||||
- 1.0
|
- 1.0
|
||||||
@@ -40,6 +40,7 @@ meta:
|
|||||||
use_bfloat16: true
|
use_bfloat16: true
|
||||||
use_gradient_checkpointing: true
|
use_gradient_checkpointing: true
|
||||||
optimization:
|
optimization:
|
||||||
|
gradient_accumulation_steps: 32
|
||||||
ema:
|
ema:
|
||||||
- 0.996
|
- 0.996
|
||||||
- 1.0
|
- 1.0
|
||||||
|
|||||||
+1
-1
@@ -167,5 +167,5 @@ def init_opt(
|
|||||||
ref_wd=wd,
|
ref_wd=wd,
|
||||||
final_wd=final_wd,
|
final_wd=final_wd,
|
||||||
T_max=int(ipe_scale*num_epochs*iterations_per_epoch))
|
T_max=int(ipe_scale*num_epochs*iterations_per_epoch))
|
||||||
scaler = torch.cuda.amp.GradScaler() if use_bfloat16 else None
|
scaler = None
|
||||||
return optimizer, scaler, scheduler, wd_scheduler
|
return optimizer, scaler, scheduler, wd_scheduler
|
||||||
|
|||||||
+21
-17
@@ -114,6 +114,7 @@ def main(args, resume_preempt=False):
|
|||||||
start_lr = args["optimization"]["start_lr"]
|
start_lr = args["optimization"]["start_lr"]
|
||||||
lr = args["optimization"]["lr"]
|
lr = args["optimization"]["lr"]
|
||||||
final_lr = args["optimization"]["final_lr"]
|
final_lr = args["optimization"]["final_lr"]
|
||||||
|
accum_steps = args["optimization"].get("gradient_accumulation_steps", 1)
|
||||||
|
|
||||||
# -- LOGGING
|
# -- LOGGING
|
||||||
folder = args["logging"]["folder"]
|
folder = args["logging"]["folder"]
|
||||||
@@ -275,6 +276,7 @@ def main(args, resume_preempt=False):
|
|||||||
torch.save(save_dict, save_path.format(epoch=f"{epoch + 1}"))
|
torch.save(save_dict, save_path.format(epoch=f"{epoch + 1}"))
|
||||||
|
|
||||||
# -- TRAINING LOOP
|
# -- TRAINING LOOP
|
||||||
|
optimizer.zero_grad()
|
||||||
for epoch in range(start_epoch, num_epochs):
|
for epoch in range(start_epoch, num_epochs):
|
||||||
logger.info("Epoch %d" % (epoch + 1))
|
logger.info("Epoch %d" % (epoch + 1))
|
||||||
|
|
||||||
@@ -300,10 +302,6 @@ def main(args, resume_preempt=False):
|
|||||||
maskB_meter.update(len(masks_pred[0][0]))
|
maskB_meter.update(len(masks_pred[0][0]))
|
||||||
|
|
||||||
def train_step():
|
def train_step():
|
||||||
_new_lr = scheduler.step()
|
|
||||||
_new_wd = wd_scheduler.step()
|
|
||||||
# --
|
|
||||||
|
|
||||||
def forward_target():
|
def forward_target():
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
h = target_encoder(imgs)
|
h = target_encoder(imgs)
|
||||||
@@ -332,24 +330,28 @@ def main(args, resume_preempt=False):
|
|||||||
z = forward_context()
|
z = forward_context()
|
||||||
loss = loss_fn(z, h)
|
loss = loss_fn(z, h)
|
||||||
|
|
||||||
# Step 2. Backward & step
|
# Step 2. Backward (accumulate gradients)
|
||||||
if use_bfloat16:
|
(loss / accum_steps).backward()
|
||||||
scaler.scale(loss).backward()
|
|
||||||
scaler.step(optimizer)
|
# Step 3. Optimizer / scheduler / momentum (every accum_steps)
|
||||||
scaler.update()
|
is_accumulated = ((itr + 1) % accum_steps == 0)
|
||||||
else:
|
if is_accumulated:
|
||||||
loss.backward()
|
_new_lr = scheduler.step()
|
||||||
optimizer.step()
|
_new_wd = wd_scheduler.step()
|
||||||
grad_stats = grad_logger(encoder.named_parameters())
|
optimizer.step()
|
||||||
optimizer.zero_grad(set_to_none=True)
|
optimizer.zero_grad(set_to_none=True)
|
||||||
|
grad_stats = grad_logger(encoder.named_parameters())
|
||||||
|
|
||||||
# Step 3. momentum update of target encoder
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
m = next(momentum_scheduler)
|
m = next(momentum_scheduler)
|
||||||
for param_q, param_k in zip(
|
for param_q, param_k in zip(
|
||||||
encoder.parameters(), target_encoder.parameters()
|
encoder.parameters(), target_encoder.parameters()
|
||||||
):
|
):
|
||||||
param_k.data.mul_(m).add_((1.0 - m) * param_q.detach().data)
|
param_k.data.mul_(m).add_((1.0 - m) * param_q.detach().data)
|
||||||
|
else:
|
||||||
|
_new_lr = None
|
||||||
|
_new_wd = None
|
||||||
|
grad_stats = None
|
||||||
|
|
||||||
return (float(loss), _new_lr, _new_wd, grad_stats)
|
return (float(loss), _new_lr, _new_wd, grad_stats)
|
||||||
|
|
||||||
@@ -363,6 +365,8 @@ def main(args, resume_preempt=False):
|
|||||||
epoch + 1, itr, loss, maskA_meter.val, maskB_meter.val, etime
|
epoch + 1, itr, loss, maskA_meter.val, maskB_meter.val, etime
|
||||||
)
|
)
|
||||||
if (itr % log_freq == 0) or np.isnan(loss) or np.isinf(loss):
|
if (itr % log_freq == 0) or np.isnan(loss) or np.isinf(loss):
|
||||||
|
lr_to_log = _new_lr if _new_lr is not None else 0.
|
||||||
|
wd_to_log = _new_wd if _new_wd is not None else 0.
|
||||||
logger.info(
|
logger.info(
|
||||||
"[%d, %5d] loss: %.3f "
|
"[%d, %5d] loss: %.3f "
|
||||||
"masks: %.1f %.1f "
|
"masks: %.1f %.1f "
|
||||||
@@ -375,8 +379,8 @@ def main(args, resume_preempt=False):
|
|||||||
loss_meter.avg,
|
loss_meter.avg,
|
||||||
maskA_meter.avg,
|
maskA_meter.avg,
|
||||||
maskB_meter.avg,
|
maskB_meter.avg,
|
||||||
_new_wd,
|
wd_to_log,
|
||||||
_new_lr,
|
lr_to_log,
|
||||||
torch.cuda.max_memory_allocated() / 1024.0**2,
|
torch.cuda.max_memory_allocated() / 1024.0**2,
|
||||||
time_meter.avg,
|
time_meter.avg,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user