integrate make_iwildcam in train.py
This commit is contained in:
@@ -8,6 +8,7 @@ logger = getLogger()
|
|||||||
def make_iwildcam(
|
def make_iwildcam(
|
||||||
transform,
|
transform,
|
||||||
batch_size,
|
batch_size,
|
||||||
|
collator=None,
|
||||||
split="extra_unlabeled",
|
split="extra_unlabeled",
|
||||||
num_workers=8,
|
num_workers=8,
|
||||||
world_size=1,
|
world_size=1,
|
||||||
@@ -36,12 +37,13 @@ def make_iwildcam(
|
|||||||
else:
|
else:
|
||||||
dataset = WildsToTorchWrapper(dataset)
|
dataset = WildsToTorchWrapper(dataset)
|
||||||
|
|
||||||
dist_sampler = torch.utils.data.DistributedSampler(
|
dist_sampler = torch.utils.data.distributed.DistributedSampler(
|
||||||
dataset=dataset, num_replicas=world_size, rank=rank, shuffle=shuffle
|
dataset=dataset, num_replicas=world_size, rank=rank, shuffle=shuffle
|
||||||
)
|
)
|
||||||
|
|
||||||
data_loader = torch.utils.data.DataLoader(
|
data_loader = torch.utils.data.DataLoader(
|
||||||
dataset,
|
dataset,
|
||||||
|
collate_fn=collator,
|
||||||
sampler=dist_sampler,
|
sampler=dist_sampler,
|
||||||
batch_size=batch_size,
|
batch_size=batch_size,
|
||||||
drop_last=drop_last,
|
drop_last=drop_last,
|
||||||
|
|||||||
+129
-111
@@ -13,7 +13,7 @@ try:
|
|||||||
# -- SURE TO UPDATE THIS TO GET LOCAL-RANK ON NODE, OR ENSURE
|
# -- SURE TO UPDATE THIS TO GET LOCAL-RANK ON NODE, OR ENSURE
|
||||||
# -- THAT YOUR JOBS ARE LAUNCHED WITH ONLY 1 DEVICE VISIBLE
|
# -- THAT YOUR JOBS ARE LAUNCHED WITH ONLY 1 DEVICE VISIBLE
|
||||||
# -- TO EACH PROCESS
|
# -- TO EACH PROCESS
|
||||||
os.environ['CUDA_VISIBLE_DEVICES'] = os.environ['SLURM_LOCALID']
|
os.environ["CUDA_VISIBLE_DEVICES"] = os.environ["SLURM_LOCALID"]
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@@ -31,22 +31,12 @@ from torch.nn.parallel import DistributedDataParallel
|
|||||||
|
|
||||||
from src.masks.multiblock import MaskCollator as MBMaskCollator
|
from src.masks.multiblock import MaskCollator as MBMaskCollator
|
||||||
from src.masks.utils import apply_masks
|
from src.masks.utils import apply_masks
|
||||||
from src.utils.distributed import (
|
from src.utils.distributed import init_distributed, AllReduce
|
||||||
init_distributed,
|
from src.utils.logging import CSVLogger, gpu_timer, grad_logger, AverageMeter
|
||||||
AllReduce
|
|
||||||
)
|
|
||||||
from src.utils.logging import (
|
|
||||||
CSVLogger,
|
|
||||||
gpu_timer,
|
|
||||||
grad_logger,
|
|
||||||
AverageMeter)
|
|
||||||
from src.utils.tensors import repeat_interleave_batch
|
from src.utils.tensors import repeat_interleave_batch
|
||||||
from src.datasets.imagenet1k import make_imagenet1k
|
from src.datasets.wilds import make_iwildcam
|
||||||
|
|
||||||
from src.helper import (
|
from src.helper import load_checkpoint, init_model, init_opt
|
||||||
load_checkpoint,
|
|
||||||
init_model,
|
|
||||||
init_opt)
|
|
||||||
from src.transforms import make_transforms
|
from src.transforms import make_transforms
|
||||||
|
|
||||||
# --
|
# --
|
||||||
@@ -65,98 +55,101 @@ logger = logging.getLogger()
|
|||||||
|
|
||||||
|
|
||||||
def main(args, resume_preempt=False):
|
def main(args, resume_preempt=False):
|
||||||
|
|
||||||
# ----------------------------------------------------------------------- #
|
# ----------------------------------------------------------------------- #
|
||||||
# PASSED IN PARAMS FROM CONFIG FILE
|
# PASSED IN PARAMS FROM CONFIG FILE
|
||||||
# ----------------------------------------------------------------------- #
|
# ----------------------------------------------------------------------- #
|
||||||
|
|
||||||
# -- META
|
# -- META
|
||||||
use_bfloat16 = args['meta']['use_bfloat16']
|
use_bfloat16 = args["meta"]["use_bfloat16"]
|
||||||
model_name = args['meta']['model_name']
|
model_name = args["meta"]["model_name"]
|
||||||
load_model = args['meta']['load_checkpoint'] or resume_preempt
|
load_model = args["meta"]["load_checkpoint"] or resume_preempt
|
||||||
r_file = args['meta']['read_checkpoint']
|
r_file = args["meta"]["read_checkpoint"]
|
||||||
copy_data = args['meta']['copy_data']
|
copy_data = args["meta"]["copy_data"]
|
||||||
pred_depth = args['meta']['pred_depth']
|
pred_depth = args["meta"]["pred_depth"]
|
||||||
pred_emb_dim = args['meta']['pred_emb_dim']
|
pred_emb_dim = args["meta"]["pred_emb_dim"]
|
||||||
if not torch.cuda.is_available():
|
if not torch.cuda.is_available():
|
||||||
device = torch.device('cpu')
|
device = torch.device("cpu")
|
||||||
else:
|
else:
|
||||||
device = torch.device('cuda:0')
|
device = torch.device("cuda:0")
|
||||||
torch.cuda.set_device(device)
|
torch.cuda.set_device(device)
|
||||||
|
|
||||||
# -- DATA
|
# -- DATA
|
||||||
use_gaussian_blur = args['data']['use_gaussian_blur']
|
use_gaussian_blur = args["data"]["use_gaussian_blur"]
|
||||||
use_horizontal_flip = args['data']['use_horizontal_flip']
|
use_horizontal_flip = args["data"]["use_horizontal_flip"]
|
||||||
use_color_distortion = args['data']['use_color_distortion']
|
use_color_distortion = args["data"]["use_color_distortion"]
|
||||||
color_jitter = args['data']['color_jitter_strength']
|
color_jitter = args["data"]["color_jitter_strength"]
|
||||||
# --
|
# --
|
||||||
batch_size = args['data']['batch_size']
|
batch_size = args["data"]["batch_size"]
|
||||||
pin_mem = args['data']['pin_mem']
|
pin_mem = args["data"]["pin_mem"]
|
||||||
num_workers = args['data']['num_workers']
|
num_workers = args["data"]["num_workers"]
|
||||||
root_path = args['data']['root_path']
|
root_path = args["data"]["root_path"]
|
||||||
image_folder = args['data']['image_folder']
|
image_folder = args["data"]["image_folder"]
|
||||||
crop_size = args['data']['crop_size']
|
crop_size = args["data"]["crop_size"]
|
||||||
crop_scale = args['data']['crop_scale']
|
crop_scale = args["data"]["crop_scale"]
|
||||||
# --
|
# --
|
||||||
|
|
||||||
# -- MASK
|
# -- MASK
|
||||||
allow_overlap = args['mask']['allow_overlap'] # whether to allow overlap b/w context and target blocks
|
allow_overlap = args["mask"][
|
||||||
patch_size = args['mask']['patch_size'] # patch-size for model training
|
"allow_overlap"
|
||||||
num_enc_masks = args['mask']['num_enc_masks'] # number of context blocks
|
] # whether to allow overlap b/w context and target blocks
|
||||||
min_keep = args['mask']['min_keep'] # min number of patches in context block
|
patch_size = args["mask"]["patch_size"] # patch-size for model training
|
||||||
enc_mask_scale = args['mask']['enc_mask_scale'] # scale of context blocks
|
num_enc_masks = args["mask"]["num_enc_masks"] # number of context blocks
|
||||||
num_pred_masks = args['mask']['num_pred_masks'] # number of target blocks
|
min_keep = args["mask"]["min_keep"] # min number of patches in context block
|
||||||
pred_mask_scale = args['mask']['pred_mask_scale'] # scale of target blocks
|
enc_mask_scale = args["mask"]["enc_mask_scale"] # scale of context blocks
|
||||||
aspect_ratio = args['mask']['aspect_ratio'] # aspect ratio of target blocks
|
num_pred_masks = args["mask"]["num_pred_masks"] # number of target blocks
|
||||||
|
pred_mask_scale = args["mask"]["pred_mask_scale"] # scale of target blocks
|
||||||
|
aspect_ratio = args["mask"]["aspect_ratio"] # aspect ratio of target blocks
|
||||||
# --
|
# --
|
||||||
|
|
||||||
# -- OPTIMIZATION
|
# -- OPTIMIZATION
|
||||||
ema = args['optimization']['ema']
|
ema = args["optimization"]["ema"]
|
||||||
ipe_scale = args['optimization']['ipe_scale'] # scheduler scale factor (def: 1.0)
|
ipe_scale = args["optimization"]["ipe_scale"] # scheduler scale factor (def: 1.0)
|
||||||
wd = float(args['optimization']['weight_decay'])
|
wd = float(args["optimization"]["weight_decay"])
|
||||||
final_wd = float(args['optimization']['final_weight_decay'])
|
final_wd = float(args["optimization"]["final_weight_decay"])
|
||||||
num_epochs = args['optimization']['epochs']
|
num_epochs = args["optimization"]["epochs"]
|
||||||
warmup = args['optimization']['warmup']
|
warmup = args["optimization"]["warmup"]
|
||||||
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"]
|
||||||
|
|
||||||
# -- LOGGING
|
# -- LOGGING
|
||||||
folder = args['logging']['folder']
|
folder = args["logging"]["folder"]
|
||||||
tag = args['logging']['write_tag']
|
tag = args["logging"]["write_tag"]
|
||||||
|
|
||||||
dump = os.path.join(folder, 'params-ijepa.yaml')
|
dump = os.path.join(folder, "params-ijepa.yaml")
|
||||||
with open(dump, 'w') as f:
|
with open(dump, "w") as f:
|
||||||
yaml.dump(args, f)
|
yaml.dump(args, f)
|
||||||
# ----------------------------------------------------------------------- #
|
# ----------------------------------------------------------------------- #
|
||||||
|
|
||||||
try:
|
try:
|
||||||
mp.set_start_method('spawn')
|
mp.set_start_method("spawn")
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
# -- init torch distributed backend
|
# -- init torch distributed backend
|
||||||
world_size, rank = init_distributed()
|
world_size, rank = init_distributed()
|
||||||
logger.info(f'Initialized (rank/world-size) {rank}/{world_size}')
|
logger.info(f"Initialized (rank/world-size) {rank}/{world_size}")
|
||||||
if rank > 0:
|
if rank > 0:
|
||||||
logger.setLevel(logging.ERROR)
|
logger.setLevel(logging.ERROR)
|
||||||
|
|
||||||
# -- log/checkpointing paths
|
# -- log/checkpointing paths
|
||||||
log_file = os.path.join(folder, f'{tag}_r{rank}.csv')
|
log_file = os.path.join(folder, f"{tag}_r{rank}.csv")
|
||||||
save_path = os.path.join(folder, f'{tag}' + '-ep{epoch}.pth.tar')
|
save_path = os.path.join(folder, f"{tag}" + "-ep{epoch}.pth.tar")
|
||||||
latest_path = os.path.join(folder, f'{tag}-latest.pth.tar')
|
latest_path = os.path.join(folder, f"{tag}-latest.pth.tar")
|
||||||
load_path = None
|
load_path = None
|
||||||
if load_model:
|
if load_model:
|
||||||
load_path = os.path.join(folder, r_file) if r_file is not None else latest_path
|
load_path = os.path.join(folder, r_file) if r_file is not None else latest_path
|
||||||
|
|
||||||
# -- make csv_logger
|
# -- make csv_logger
|
||||||
csv_logger = CSVLogger(log_file,
|
csv_logger = CSVLogger(
|
||||||
('%d', 'epoch'),
|
log_file,
|
||||||
('%d', 'itr'),
|
("%d", "epoch"),
|
||||||
('%.5f', 'loss'),
|
("%d", "itr"),
|
||||||
('%.5f', 'mask-A'),
|
("%.5f", "loss"),
|
||||||
('%.5f', 'mask-B'),
|
("%.5f", "mask-A"),
|
||||||
('%d', 'time (ms)'))
|
("%.5f", "mask-B"),
|
||||||
|
("%d", "time (ms)"),
|
||||||
|
)
|
||||||
|
|
||||||
# -- init model
|
# -- init model
|
||||||
encoder, predictor = init_model(
|
encoder, predictor = init_model(
|
||||||
@@ -165,7 +158,8 @@ def main(args, resume_preempt=False):
|
|||||||
crop_size=crop_size,
|
crop_size=crop_size,
|
||||||
pred_depth=pred_depth,
|
pred_depth=pred_depth,
|
||||||
pred_emb_dim=pred_emb_dim,
|
pred_emb_dim=pred_emb_dim,
|
||||||
model_name=model_name)
|
model_name=model_name,
|
||||||
|
)
|
||||||
target_encoder = copy.deepcopy(encoder)
|
target_encoder = copy.deepcopy(encoder)
|
||||||
|
|
||||||
# -- make data transforms
|
# -- make data transforms
|
||||||
@@ -178,7 +172,8 @@ def main(args, resume_preempt=False):
|
|||||||
nenc=num_enc_masks,
|
nenc=num_enc_masks,
|
||||||
npred=num_pred_masks,
|
npred=num_pred_masks,
|
||||||
allow_overlap=allow_overlap,
|
allow_overlap=allow_overlap,
|
||||||
min_keep=min_keep)
|
min_keep=min_keep,
|
||||||
|
)
|
||||||
|
|
||||||
transform = make_transforms(
|
transform = make_transforms(
|
||||||
crop_size=crop_size,
|
crop_size=crop_size,
|
||||||
@@ -186,22 +181,21 @@ def main(args, resume_preempt=False):
|
|||||||
gaussian_blur=use_gaussian_blur,
|
gaussian_blur=use_gaussian_blur,
|
||||||
horizontal_flip=use_horizontal_flip,
|
horizontal_flip=use_horizontal_flip,
|
||||||
color_distortion=use_color_distortion,
|
color_distortion=use_color_distortion,
|
||||||
color_jitter=color_jitter)
|
color_jitter=color_jitter,
|
||||||
|
)
|
||||||
|
|
||||||
# -- init data-loaders/samplers
|
# -- init data-loaders/samplers
|
||||||
_, unsupervised_loader, unsupervised_sampler = make_imagenet1k(
|
_, unsupervised_loader, unsupervised_sampler = make_iwildcam(
|
||||||
transform=transform,
|
transform=transform,
|
||||||
batch_size=batch_size,
|
batch_size=batch_size,
|
||||||
collator=mask_collator,
|
collator=mask_collator,
|
||||||
pin_mem=pin_mem,
|
pin_mem=pin_mem,
|
||||||
training=True,
|
|
||||||
num_workers=num_workers,
|
num_workers=num_workers,
|
||||||
world_size=world_size,
|
world_size=world_size,
|
||||||
rank=rank,
|
rank=rank,
|
||||||
root_path=root_path,
|
root_path=root_path,
|
||||||
image_folder=image_folder,
|
drop_last=True,
|
||||||
copy_data=copy_data,
|
)
|
||||||
drop_last=True)
|
|
||||||
ipe = len(unsupervised_loader)
|
ipe = len(unsupervised_loader)
|
||||||
|
|
||||||
# -- init optimizer and scheduler
|
# -- init optimizer and scheduler
|
||||||
@@ -217,7 +211,8 @@ def main(args, resume_preempt=False):
|
|||||||
warmup=warmup,
|
warmup=warmup,
|
||||||
num_epochs=num_epochs,
|
num_epochs=num_epochs,
|
||||||
ipe_scale=ipe_scale,
|
ipe_scale=ipe_scale,
|
||||||
use_bfloat16=use_bfloat16)
|
use_bfloat16=use_bfloat16,
|
||||||
|
)
|
||||||
encoder = DistributedDataParallel(encoder, static_graph=True)
|
encoder = DistributedDataParallel(encoder, static_graph=True)
|
||||||
predictor = DistributedDataParallel(predictor, static_graph=True)
|
predictor = DistributedDataParallel(predictor, static_graph=True)
|
||||||
target_encoder = DistributedDataParallel(target_encoder)
|
target_encoder = DistributedDataParallel(target_encoder)
|
||||||
@@ -225,20 +220,25 @@ def main(args, resume_preempt=False):
|
|||||||
p.requires_grad = False
|
p.requires_grad = False
|
||||||
|
|
||||||
# -- momentum schedule
|
# -- momentum schedule
|
||||||
momentum_scheduler = (ema[0] + i*(ema[1]-ema[0])/(ipe*num_epochs*ipe_scale)
|
momentum_scheduler = (
|
||||||
for i in range(int(ipe*num_epochs*ipe_scale)+1))
|
ema[0] + i * (ema[1] - ema[0]) / (ipe * num_epochs * ipe_scale)
|
||||||
|
for i in range(int(ipe * num_epochs * ipe_scale) + 1)
|
||||||
|
)
|
||||||
|
|
||||||
start_epoch = 0
|
start_epoch = 0
|
||||||
# -- load training checkpoint
|
# -- load training checkpoint
|
||||||
if load_model:
|
if load_model:
|
||||||
encoder, predictor, target_encoder, optimizer, scaler, start_epoch = load_checkpoint(
|
encoder, predictor, target_encoder, optimizer, scaler, start_epoch = (
|
||||||
|
load_checkpoint(
|
||||||
device=device,
|
device=device,
|
||||||
r_path=load_path,
|
r_path=load_path,
|
||||||
encoder=encoder,
|
encoder=encoder,
|
||||||
predictor=predictor,
|
predictor=predictor,
|
||||||
target_encoder=target_encoder,
|
target_encoder=target_encoder,
|
||||||
opt=optimizer,
|
opt=optimizer,
|
||||||
scaler=scaler)
|
scaler=scaler,
|
||||||
|
)
|
||||||
|
)
|
||||||
for _ in range(start_epoch * ipe):
|
for _ in range(start_epoch * ipe):
|
||||||
scheduler.step()
|
scheduler.step()
|
||||||
wd_scheduler.step()
|
wd_scheduler.step()
|
||||||
@@ -247,25 +247,25 @@ def main(args, resume_preempt=False):
|
|||||||
|
|
||||||
def save_checkpoint(epoch):
|
def save_checkpoint(epoch):
|
||||||
save_dict = {
|
save_dict = {
|
||||||
'encoder': encoder.state_dict(),
|
"encoder": encoder.state_dict(),
|
||||||
'predictor': predictor.state_dict(),
|
"predictor": predictor.state_dict(),
|
||||||
'target_encoder': target_encoder.state_dict(),
|
"target_encoder": target_encoder.state_dict(),
|
||||||
'opt': optimizer.state_dict(),
|
"opt": optimizer.state_dict(),
|
||||||
'scaler': None if scaler is None else scaler.state_dict(),
|
"scaler": None if scaler is None else scaler.state_dict(),
|
||||||
'epoch': epoch,
|
"epoch": epoch,
|
||||||
'loss': loss_meter.avg,
|
"loss": loss_meter.avg,
|
||||||
'batch_size': batch_size,
|
"batch_size": batch_size,
|
||||||
'world_size': world_size,
|
"world_size": world_size,
|
||||||
'lr': lr
|
"lr": lr,
|
||||||
}
|
}
|
||||||
if rank == 0:
|
if rank == 0:
|
||||||
torch.save(save_dict, latest_path)
|
torch.save(save_dict, latest_path)
|
||||||
if (epoch + 1) % checkpoint_freq == 0:
|
if (epoch + 1) % checkpoint_freq == 0:
|
||||||
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
|
||||||
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))
|
||||||
|
|
||||||
# -- update distributed-data-loader epoch
|
# -- update distributed-data-loader epoch
|
||||||
unsupervised_sampler.set_epoch(epoch)
|
unsupervised_sampler.set_epoch(epoch)
|
||||||
@@ -283,6 +283,7 @@ def main(args, resume_preempt=False):
|
|||||||
masks_1 = [u.to(device, non_blocking=True) for u in masks_enc]
|
masks_1 = [u.to(device, non_blocking=True) for u in masks_enc]
|
||||||
masks_2 = [u.to(device, non_blocking=True) for u in masks_pred]
|
masks_2 = [u.to(device, non_blocking=True) for u in masks_pred]
|
||||||
return (imgs, masks_1, masks_2)
|
return (imgs, masks_1, masks_2)
|
||||||
|
|
||||||
imgs, masks_enc, masks_pred = load_imgs()
|
imgs, masks_enc, masks_pred = load_imgs()
|
||||||
maskA_meter.update(len(masks_enc[0][0]))
|
maskA_meter.update(len(masks_enc[0][0]))
|
||||||
maskB_meter.update(len(masks_pred[0][0]))
|
maskB_meter.update(len(masks_pred[0][0]))
|
||||||
@@ -313,7 +314,9 @@ def main(args, resume_preempt=False):
|
|||||||
return loss
|
return loss
|
||||||
|
|
||||||
# Step 1. Forward
|
# Step 1. Forward
|
||||||
with torch.cuda.amp.autocast(dtype=torch.bfloat16, enabled=use_bfloat16):
|
with torch.cuda.amp.autocast(
|
||||||
|
dtype=torch.bfloat16, enabled=use_bfloat16
|
||||||
|
):
|
||||||
h = forward_target()
|
h = forward_target()
|
||||||
z = forward_context()
|
z = forward_context()
|
||||||
loss = loss_fn(z, h)
|
loss = loss_fn(z, h)
|
||||||
@@ -332,46 +335,61 @@ def main(args, resume_preempt=False):
|
|||||||
# Step 3. momentum update of target encoder
|
# 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(encoder.parameters(), target_encoder.parameters()):
|
for param_q, param_k in zip(
|
||||||
param_k.data.mul_(m).add_((1.-m) * param_q.detach().data)
|
encoder.parameters(), target_encoder.parameters()
|
||||||
|
):
|
||||||
|
param_k.data.mul_(m).add_((1.0 - m) * param_q.detach().data)
|
||||||
|
|
||||||
return (float(loss), _new_lr, _new_wd, grad_stats)
|
return (float(loss), _new_lr, _new_wd, grad_stats)
|
||||||
|
|
||||||
(loss, _new_lr, _new_wd, grad_stats), etime = gpu_timer(train_step)
|
(loss, _new_lr, _new_wd, grad_stats), etime = gpu_timer(train_step)
|
||||||
loss_meter.update(loss)
|
loss_meter.update(loss)
|
||||||
time_meter.update(etime)
|
time_meter.update(etime)
|
||||||
|
|
||||||
# -- Logging
|
# -- Logging
|
||||||
def log_stats():
|
def log_stats():
|
||||||
csv_logger.log(epoch + 1, itr, loss, maskA_meter.val, maskB_meter.val, etime)
|
csv_logger.log(
|
||||||
|
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):
|
||||||
logger.info('[%d, %5d] loss: %.3f '
|
logger.info(
|
||||||
'masks: %.1f %.1f '
|
"[%d, %5d] loss: %.3f "
|
||||||
'[wd: %.2e] [lr: %.2e] '
|
"masks: %.1f %.1f "
|
||||||
'[mem: %.2e] '
|
"[wd: %.2e] [lr: %.2e] "
|
||||||
'(%.1f ms)'
|
"[mem: %.2e] "
|
||||||
% (epoch + 1, itr,
|
"(%.1f ms)"
|
||||||
|
% (
|
||||||
|
epoch + 1,
|
||||||
|
itr,
|
||||||
loss_meter.avg,
|
loss_meter.avg,
|
||||||
maskA_meter.avg,
|
maskA_meter.avg,
|
||||||
maskB_meter.avg,
|
maskB_meter.avg,
|
||||||
_new_wd,
|
_new_wd,
|
||||||
_new_lr,
|
_new_lr,
|
||||||
torch.cuda.max_memory_allocated() / 1024.**2,
|
torch.cuda.max_memory_allocated() / 1024.0**2,
|
||||||
time_meter.avg))
|
time_meter.avg,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
if grad_stats is not None:
|
if grad_stats is not None:
|
||||||
logger.info('[%d, %5d] grad_stats: [%.2e %.2e] (%.2e, %.2e)'
|
logger.info(
|
||||||
% (epoch + 1, itr,
|
"[%d, %5d] grad_stats: [%.2e %.2e] (%.2e, %.2e)"
|
||||||
|
% (
|
||||||
|
epoch + 1,
|
||||||
|
itr,
|
||||||
grad_stats.first_layer,
|
grad_stats.first_layer,
|
||||||
grad_stats.last_layer,
|
grad_stats.last_layer,
|
||||||
grad_stats.min,
|
grad_stats.min,
|
||||||
grad_stats.max))
|
grad_stats.max,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
log_stats()
|
log_stats()
|
||||||
|
|
||||||
assert not np.isnan(loss), 'loss is nan'
|
assert not np.isnan(loss), "loss is nan"
|
||||||
|
|
||||||
# -- Save Checkpoint after every epoch
|
# -- Save Checkpoint after every epoch
|
||||||
logger.info('avg. loss %.3f' % loss_meter.avg)
|
logger.info("avg. loss %.3f" % loss_meter.avg)
|
||||||
save_checkpoint(epoch + 1)
|
save_checkpoint(epoch + 1)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user