Merge branch 'main' of github.com:YannAhlgrim/ijepa
This commit is contained in:
@@ -16,13 +16,23 @@ data:
|
||||
optimization:
|
||||
optimizer: adamw # 'adamw' or 'sgd'
|
||||
freeze_weights: true # true for linear probing, false for full fine-tuning
|
||||
epochs: 50
|
||||
epochs: 200 # can be set higher if early_stopping
|
||||
lr: 0.001
|
||||
weight_decay: 0.05
|
||||
use_cosine_schedule: false
|
||||
start_lr: 0.0001
|
||||
final_lr: 1.0e-06
|
||||
warmup: 5
|
||||
ipe_scale: 1.0
|
||||
early_stopping:
|
||||
enabled: true
|
||||
patience: 6
|
||||
min_delta: 0.0
|
||||
min_epochs: 5
|
||||
restore_best_weights: true
|
||||
|
||||
validation:
|
||||
eval_every: 1
|
||||
|
||||
logging:
|
||||
folder: experiment_logs/supervised-vith14.224-bs.128-ep.300/
|
||||
|
||||
@@ -31,7 +31,7 @@ mask:
|
||||
- 0.2
|
||||
meta:
|
||||
copy_data: false
|
||||
load_checkpoint: false
|
||||
load_checkpoint: true
|
||||
model_name: vit_huge
|
||||
pred_depth: 12
|
||||
pred_emb_dim: 384
|
||||
|
||||
+5
-1
@@ -13,7 +13,11 @@ class ViTClassifier(nn.Module):
|
||||
|
||||
def forward(self, x):
|
||||
# ViT -> (B, N, D)
|
||||
features = self.encoder(x)
|
||||
if any(p.requires_grad for p in self.encoder.parameters()):
|
||||
features = self.encoder(x)
|
||||
else:
|
||||
with torch.no_grad():
|
||||
features = self.encoder(x)
|
||||
|
||||
# Average Pool -> (B, D)
|
||||
avg_embed = features.mean(dim=1)
|
||||
|
||||
+337
-68
@@ -2,21 +2,22 @@ import os
|
||||
import sys
|
||||
import yaml
|
||||
import logging
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn.parallel import DistributedDataParallel
|
||||
import torch.nn.functional as F
|
||||
|
||||
import numpy as np
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.distributed as dist
|
||||
from torch.nn.parallel import DistributedDataParallel
|
||||
|
||||
from src.datasets.wilds import make_iwildcam
|
||||
from src.helper import load_checkpoint, init_model, init_opt
|
||||
from src.transforms import make_transforms
|
||||
from src.models.head import ViTClassifier # Import our new class
|
||||
from src.helper import init_model
|
||||
from src.models.head import ViTClassifier
|
||||
from src.transforms import make_transforms, make_transform_eval
|
||||
from src.utils.distributed import init_distributed
|
||||
from src.utils.logging import CSVLogger, AverageMeter
|
||||
|
||||
# --
|
||||
log_timings = True
|
||||
log_freq = 10
|
||||
checkpoint_freq = 50
|
||||
# --
|
||||
@@ -30,27 +31,172 @@ logging.basicConfig(stream=sys.stdout, level=logging.INFO)
|
||||
logger = logging.getLogger()
|
||||
|
||||
|
||||
def main(args):
|
||||
# -- Init Distributed
|
||||
def strip_module_prefix(state_dict):
|
||||
if not any(k.startswith("module.") for k in state_dict.keys()):
|
||||
return state_dict
|
||||
return {
|
||||
k[len("module.") :] if k.startswith("module.") else k: v
|
||||
for k, v in state_dict.items()
|
||||
}
|
||||
|
||||
|
||||
def distributed_average(value, device):
|
||||
tensor = torch.tensor([value], device=device, dtype=torch.float64)
|
||||
if dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1:
|
||||
dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
|
||||
tensor /= dist.get_world_size()
|
||||
return float(tensor.item())
|
||||
|
||||
|
||||
def distributed_sum(value, device):
|
||||
tensor = torch.tensor([value], device=device, dtype=torch.float64)
|
||||
if dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1:
|
||||
dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
|
||||
return float(tensor.item())
|
||||
|
||||
|
||||
class EarlyStopping:
|
||||
def __init__(
|
||||
self,
|
||||
enabled=True,
|
||||
patience=10,
|
||||
min_delta=0.0,
|
||||
min_epochs=0,
|
||||
restore_best_weights=True,
|
||||
):
|
||||
self.enabled = enabled
|
||||
self.patience = patience
|
||||
self.min_delta = min_delta
|
||||
self.min_epochs = min_epochs
|
||||
self.restore_best_weights = restore_best_weights
|
||||
self.best_metric = float("inf")
|
||||
self.best_epoch = -1
|
||||
self.bad_epochs = 0
|
||||
self.best_state = None
|
||||
|
||||
def state_dict(self):
|
||||
return {
|
||||
"enabled": self.enabled,
|
||||
"patience": self.patience,
|
||||
"min_delta": self.min_delta,
|
||||
"min_epochs": self.min_epochs,
|
||||
"restore_best_weights": self.restore_best_weights,
|
||||
"best_metric": self.best_metric,
|
||||
"best_epoch": self.best_epoch,
|
||||
"bad_epochs": self.bad_epochs,
|
||||
}
|
||||
|
||||
def load_state_dict(self, state):
|
||||
if not state:
|
||||
return
|
||||
self.enabled = state.get("enabled", self.enabled)
|
||||
self.patience = state.get("patience", self.patience)
|
||||
self.min_delta = state.get("min_delta", self.min_delta)
|
||||
self.min_epochs = state.get("min_epochs", self.min_epochs)
|
||||
self.restore_best_weights = state.get(
|
||||
"restore_best_weights", self.restore_best_weights
|
||||
)
|
||||
self.best_metric = state.get("best_metric", self.best_metric)
|
||||
self.best_epoch = state.get("best_epoch", self.best_epoch)
|
||||
self.bad_epochs = state.get("bad_epochs", self.bad_epochs)
|
||||
|
||||
def step(self, epoch, metric, model_module):
|
||||
if not self.enabled:
|
||||
return False, False
|
||||
|
||||
improved = metric < (self.best_metric - self.min_delta)
|
||||
if improved:
|
||||
self.best_metric = metric
|
||||
self.best_epoch = epoch
|
||||
self.bad_epochs = 0
|
||||
if self.restore_best_weights:
|
||||
self.best_state = {
|
||||
k: v.detach().cpu().clone()
|
||||
for k, v in model_module.state_dict().items()
|
||||
}
|
||||
return True, False
|
||||
|
||||
self.bad_epochs += 1
|
||||
should_stop = (
|
||||
epoch + 1
|
||||
) >= self.min_epochs and self.bad_epochs >= self.patience
|
||||
return False, should_stop
|
||||
|
||||
def restore(self, model_module, device):
|
||||
if self.restore_best_weights and self.best_state is not None:
|
||||
model_module.load_state_dict(self.best_state)
|
||||
model_module.to(device)
|
||||
|
||||
|
||||
def evaluate(model, loader, criterion, device, use_bfloat16):
|
||||
model.eval()
|
||||
loss_sum = 0.0
|
||||
n_correct = 0.0
|
||||
n_total = 0.0
|
||||
|
||||
with torch.no_grad():
|
||||
for imgs, labels in loader:
|
||||
imgs = imgs.to(device, non_blocking=True)
|
||||
labels = labels.to(device, non_blocking=True)
|
||||
|
||||
with torch.cuda.amp.autocast(enabled=use_bfloat16, dtype=torch.bfloat16):
|
||||
outputs = model(imgs)
|
||||
loss = criterion(outputs, labels)
|
||||
|
||||
batch_size = labels.size(0)
|
||||
preds = outputs.argmax(dim=1)
|
||||
n_correct += float((preds == labels).sum().item())
|
||||
n_total += float(batch_size)
|
||||
loss_sum += float(loss.item()) * float(batch_size)
|
||||
|
||||
global_loss_sum = distributed_sum(loss_sum, device)
|
||||
global_correct = distributed_sum(n_correct, device)
|
||||
global_total = distributed_sum(n_total, device)
|
||||
|
||||
val_loss = global_loss_sum / max(global_total, 1.0)
|
||||
val_acc = global_correct / max(global_total, 1.0)
|
||||
return val_loss, val_acc
|
||||
|
||||
|
||||
def main(args, resume_preempt=False):
|
||||
del resume_preempt
|
||||
|
||||
world_size, rank = init_distributed()
|
||||
if not torch.cuda.is_available():
|
||||
raise RuntimeError("CUDA is required for supervised distributed training")
|
||||
device = torch.device(f"cuda:{torch.cuda.current_device()}")
|
||||
|
||||
# -- Extract Config Params
|
||||
m_args = args["meta"]
|
||||
o_args = args["optimization"]
|
||||
d_args = args["data"]
|
||||
l_args = args["logging"]
|
||||
v_args = args.get("validation", {})
|
||||
es_args = o_args.get("early_stopping", {})
|
||||
|
||||
# -- Paths for saving
|
||||
folder = l_args["folder"]
|
||||
tag = l_args["write_tag"]
|
||||
if not os.path.exists(folder):
|
||||
os.makedirs(folder, exist_ok=True)
|
||||
os.makedirs(folder, exist_ok=True)
|
||||
|
||||
with open(os.path.join(folder, "params-supervised.yaml"), "w") as f:
|
||||
yaml.dump(args, f)
|
||||
|
||||
save_path = os.path.join(folder, f"{tag}" + "-ep{epoch}.pth.tar")
|
||||
latest_path = os.path.join(folder, f"{tag}-latest.pth.tar")
|
||||
best_path = os.path.join(folder, f"{tag}-best.pth.tar")
|
||||
log_file = os.path.join(folder, f"{tag}_r{rank}.csv")
|
||||
|
||||
csv_logger = CSVLogger(
|
||||
log_file,
|
||||
("%d", "epoch"),
|
||||
("%.6f", "train_loss"),
|
||||
("%.6f", "val_loss"),
|
||||
("%.6f", "val_acc"),
|
||||
("%.6e", "lr"),
|
||||
("%.6f", "best_val_loss"),
|
||||
("%d", "best_epoch"),
|
||||
("%d", "early_stop"),
|
||||
)
|
||||
|
||||
# -- 1. Initialize Encoder
|
||||
encoder, _ = init_model(
|
||||
device=device,
|
||||
patch_size=args.get("mask", {}).get("patch_size", 14),
|
||||
@@ -58,19 +204,17 @@ def main(args):
|
||||
model_name=m_args["model_name"],
|
||||
)
|
||||
|
||||
# -- 2. Wrap in Classification Head
|
||||
embed_dim = m_args.get("embed_dim")
|
||||
model = ViTClassifier(encoder, m_args["num_classes"], embed_dim).to(device)
|
||||
|
||||
# -- 3. Handle Freezing
|
||||
if o_args["freeze_weights"]:
|
||||
logger.info("Freezing encoder weights (Linear Probing mode)")
|
||||
for name, param in model.encoder.named_parameters():
|
||||
for param in model.encoder.parameters():
|
||||
param.requires_grad = False
|
||||
model.encoder.eval()
|
||||
else:
|
||||
logger.info("Training full model (Fine-tuning mode)")
|
||||
|
||||
# -- 4. Optimizer Selection
|
||||
params = [p for p in model.parameters() if p.requires_grad]
|
||||
if o_args["optimizer"].lower() == "adamw":
|
||||
optimizer = torch.optim.AdamW(
|
||||
@@ -81,76 +225,138 @@ def main(args):
|
||||
params, lr=o_args["lr"], momentum=0.9, weight_decay=o_args["weight_decay"]
|
||||
)
|
||||
|
||||
# -- 5. Resume/Load Logic
|
||||
start_epoch = 0
|
||||
# Priority 1: Check if we are resuming from an interrupted run (latest-path)
|
||||
# Priority 2: Check if we are loading a specific pre-trained checkpoint (m_args["load_checkpoint"])
|
||||
scheduler = None
|
||||
if o_args.get("use_cosine_schedule", False):
|
||||
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
|
||||
optimizer, T_max=o_args["epochs"], eta_min=o_args.get("final_lr", 0.0)
|
||||
)
|
||||
|
||||
checkpoint_to_load = None
|
||||
resuming_interrupted = False
|
||||
train_transform = make_transforms(
|
||||
crop_size=d_args["crop_size"],
|
||||
crop_scale=tuple(d_args.get("crop_scale", (0.3, 1.0))),
|
||||
horizontal_flip=d_args.get("use_horizontal_flip", False),
|
||||
color_distortion=d_args.get("use_color_distortion", False),
|
||||
color_jitter=d_args.get("color_jitter_strength", 1.0),
|
||||
gaussian_blur=d_args.get("use_gaussian_blur", False),
|
||||
)
|
||||
val_transform = make_transform_eval(
|
||||
crop_size=d_args["crop_size"],
|
||||
)
|
||||
|
||||
if os.path.exists(latest_path):
|
||||
checkpoint_to_load = latest_path
|
||||
resuming_interrupted = True
|
||||
elif m_args["load_checkpoint"]:
|
||||
checkpoint_to_load = os.path.join(folder, m_args["read_checkpoint"])
|
||||
|
||||
if checkpoint_to_load:
|
||||
checkpoint = torch.load(checkpoint_to_load, map_location="cpu")
|
||||
|
||||
if resuming_interrupted:
|
||||
# Load full state to resume exactly where we left off
|
||||
model.load_state_dict(checkpoint["model"])
|
||||
optimizer.load_state_dict(checkpoint["opt"])
|
||||
start_epoch = checkpoint["epoch"]
|
||||
logger.info(
|
||||
f"Resuming training from {checkpoint_to_load} at epoch {start_epoch}"
|
||||
)
|
||||
else:
|
||||
# Loading just encoder weights for a fresh supervised run
|
||||
msg = model.encoder.load_state_dict(checkpoint["encoder"], strict=False)
|
||||
logger.info(
|
||||
f"Loaded pre-trained encoder from {checkpoint_to_load} with msg: {msg}"
|
||||
)
|
||||
|
||||
# -- 6. Data Setup
|
||||
transform = make_transforms(crop_size=d_args["crop_size"])
|
||||
_, loader, sampler = make_iwildcam(
|
||||
transform=transform,
|
||||
_, train_loader, train_sampler = make_iwildcam(
|
||||
transform=train_transform,
|
||||
split="train",
|
||||
batch_size=d_args["batch_size"],
|
||||
root_path=d_args["root_path"],
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
collator=None,
|
||||
num_workers=d_args.get("num_workers", 8),
|
||||
pin_mem=d_args.get("pin_mem", True),
|
||||
drop_last=True,
|
||||
)
|
||||
|
||||
_, val_loader, val_sampler = make_iwildcam(
|
||||
transform=val_transform,
|
||||
split="val",
|
||||
batch_size=d_args["batch_size"],
|
||||
root_path=d_args["root_path"],
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
collator=None,
|
||||
num_workers=d_args.get("num_workers", 8),
|
||||
pin_mem=d_args.get("pin_mem", True),
|
||||
drop_last=False,
|
||||
)
|
||||
|
||||
criterion = nn.CrossEntropyLoss().to(device)
|
||||
model = DistributedDataParallel(model, device_ids=[torch.cuda.current_device()])
|
||||
|
||||
# -- 7. Define Save Function
|
||||
def save_checkpoint(epoch, current_loss):
|
||||
early_stopper = EarlyStopping(
|
||||
enabled=es_args.get("enabled", False),
|
||||
patience=es_args.get("patience", 10),
|
||||
min_delta=es_args.get("min_delta", 0.0),
|
||||
min_epochs=es_args.get("min_epochs", 0),
|
||||
restore_best_weights=es_args.get("restore_best_weights", True),
|
||||
)
|
||||
|
||||
start_epoch = 0
|
||||
|
||||
checkpoint_to_load = None
|
||||
resuming_interrupted = False
|
||||
if os.path.exists(latest_path):
|
||||
checkpoint_to_load = latest_path
|
||||
resuming_interrupted = True
|
||||
elif m_args.get("load_checkpoint", False):
|
||||
r_file = m_args.get("read_checkpoint")
|
||||
if r_file is not None:
|
||||
checkpoint_to_load = os.path.join(folder, r_file)
|
||||
|
||||
if checkpoint_to_load and os.path.exists(checkpoint_to_load):
|
||||
checkpoint = torch.load(checkpoint_to_load, map_location="cpu")
|
||||
if resuming_interrupted and "model" in checkpoint:
|
||||
model.module.load_state_dict(checkpoint["model"])
|
||||
if "opt" in checkpoint:
|
||||
optimizer.load_state_dict(checkpoint["opt"])
|
||||
if scheduler is not None and "scheduler" in checkpoint:
|
||||
scheduler.load_state_dict(checkpoint["scheduler"])
|
||||
start_epoch = int(checkpoint.get("epoch", 0))
|
||||
early_stopper.load_state_dict(checkpoint.get("early_stopping", {}))
|
||||
logger.info(
|
||||
f"Resuming training from {checkpoint_to_load} at epoch {start_epoch}"
|
||||
)
|
||||
else:
|
||||
encoder_state = checkpoint.get("encoder")
|
||||
if encoder_state is None and "model" in checkpoint:
|
||||
encoder_state = {
|
||||
k.replace("encoder.", "", 1): v
|
||||
for k, v in checkpoint["model"].items()
|
||||
if k.startswith("encoder.")
|
||||
}
|
||||
if encoder_state is None:
|
||||
raise KeyError(
|
||||
f"No encoder weights found in checkpoint: {checkpoint_to_load}"
|
||||
)
|
||||
encoder_state = strip_module_prefix(encoder_state)
|
||||
msg = model.module.encoder.load_state_dict(encoder_state, strict=False)
|
||||
logger.info(
|
||||
f"Loaded pre-trained encoder from {checkpoint_to_load} with msg: {msg}"
|
||||
)
|
||||
|
||||
def save_checkpoint(epoch, train_loss, val_loss, val_acc, is_best=False):
|
||||
save_dict = {
|
||||
"model": model.module.state_dict(), # model.module because of DDP
|
||||
"model": model.module.state_dict(),
|
||||
"opt": optimizer.state_dict(),
|
||||
"scheduler": None if scheduler is None else scheduler.state_dict(),
|
||||
"epoch": epoch,
|
||||
"loss": current_loss,
|
||||
"train_loss": train_loss,
|
||||
"val_loss": val_loss,
|
||||
"val_acc": val_acc,
|
||||
"args": args,
|
||||
"early_stopping": early_stopper.state_dict(),
|
||||
}
|
||||
if rank == 0:
|
||||
torch.save(save_dict, latest_path)
|
||||
if epoch % checkpoint_freq == 0:
|
||||
torch.save(save_dict, save_path.format(epoch=epoch))
|
||||
logger.info(f"Checkpoint saved at epoch {epoch}")
|
||||
if is_best:
|
||||
torch.save(save_dict, best_path)
|
||||
|
||||
eval_every = int(v_args.get("eval_every", 1))
|
||||
|
||||
# -- 8. Training Loop
|
||||
for epoch in range(start_epoch, o_args["epochs"]):
|
||||
sampler.set_epoch(epoch)
|
||||
train_sampler.set_epoch(epoch)
|
||||
val_sampler.set_epoch(epoch)
|
||||
|
||||
model.train()
|
||||
if o_args["freeze_weights"]:
|
||||
model.module.encoder.eval()
|
||||
|
||||
loss_meter = AverageMeter()
|
||||
|
||||
for itr, (imgs, labels, _) in enumerate(loader):
|
||||
imgs, labels = imgs.to(device), labels.to(device)
|
||||
for itr, (imgs, labels) in enumerate(train_loader):
|
||||
imgs = imgs.to(device, non_blocking=True)
|
||||
labels = labels.to(device, non_blocking=True)
|
||||
|
||||
with torch.cuda.amp.autocast(
|
||||
enabled=m_args["use_bfloat16"], dtype=torch.bfloat16
|
||||
@@ -162,16 +368,79 @@ def main(args):
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
loss_meter.update(loss.item())
|
||||
loss_meter.update(loss.item(), n=labels.size(0))
|
||||
|
||||
if itr % 10 == 0 and rank == 0:
|
||||
if itr % log_freq == 0 and rank == 0:
|
||||
logger.info(
|
||||
f"Epoch {epoch + 1} [{itr}/{len(loader)}] Loss: {loss_meter.avg:.4f}"
|
||||
f"Epoch {epoch + 1} [{itr}/{len(train_loader)}] Train Loss: {loss_meter.avg:.4f}"
|
||||
)
|
||||
|
||||
# Save at the end of every epoch
|
||||
save_checkpoint(epoch + 1, loss_meter.avg)
|
||||
train_loss = distributed_average(loss_meter.avg, device)
|
||||
|
||||
do_eval = ((epoch + 1) % eval_every == 0) or (epoch + 1 == o_args["epochs"])
|
||||
if do_eval:
|
||||
val_loss, val_acc = evaluate(
|
||||
model=model,
|
||||
loader=val_loader,
|
||||
criterion=criterion,
|
||||
device=device,
|
||||
use_bfloat16=m_args["use_bfloat16"],
|
||||
)
|
||||
is_best, should_stop = early_stopper.step(epoch + 1, val_loss, model.module)
|
||||
else:
|
||||
val_loss = float("nan")
|
||||
val_acc = float("nan")
|
||||
is_best, should_stop = False, False
|
||||
|
||||
if scheduler is not None:
|
||||
scheduler.step()
|
||||
|
||||
if rank == 0:
|
||||
logger.info(
|
||||
f"Epoch {epoch + 1} done | train_loss={train_loss:.6f} val_loss={val_loss:.6f} val_acc={val_acc:.6f} best_val_loss={early_stopper.best_metric:.6f}"
|
||||
)
|
||||
|
||||
csv_logger.log(
|
||||
epoch + 1,
|
||||
train_loss,
|
||||
val_loss,
|
||||
val_acc,
|
||||
optimizer.param_groups[0]["lr"],
|
||||
early_stopper.best_metric,
|
||||
early_stopper.best_epoch,
|
||||
int(should_stop),
|
||||
)
|
||||
|
||||
save_checkpoint(epoch + 1, train_loss, val_loss, val_acc, is_best=is_best)
|
||||
|
||||
stop_tensor = torch.tensor([int(should_stop)], device=device)
|
||||
if dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1:
|
||||
dist.broadcast(stop_tensor, src=0)
|
||||
|
||||
if bool(stop_tensor.item()):
|
||||
if rank == 0:
|
||||
logger.info(
|
||||
f"Early stopping at epoch {epoch + 1}. Best val_loss={early_stopper.best_metric:.6f} @ epoch {early_stopper.best_epoch}"
|
||||
)
|
||||
break
|
||||
|
||||
if early_stopper.enabled and early_stopper.restore_best_weights:
|
||||
if rank == 0:
|
||||
logger.info("Restoring best model weights before exit")
|
||||
early_stopper.restore(model.module, device)
|
||||
if rank == 0:
|
||||
torch.save(
|
||||
{
|
||||
"model": model.module.state_dict(),
|
||||
"epoch": early_stopper.best_epoch,
|
||||
"val_loss": early_stopper.best_metric,
|
||||
"args": args,
|
||||
},
|
||||
best_path,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
raise RuntimeError(
|
||||
"Use main_distributed_supervised.py to launch this script with a config file."
|
||||
)
|
||||
|
||||
+22
-8
@@ -23,19 +23,16 @@ def make_transforms(
|
||||
horizontal_flip=False,
|
||||
color_distortion=False,
|
||||
gaussian_blur=False,
|
||||
normalization=((0.485, 0.456, 0.406),
|
||||
(0.229, 0.224, 0.225))
|
||||
normalization=((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),
|
||||
):
|
||||
logger.info('making imagenet data transforms')
|
||||
logger.info("making imagenet data transforms")
|
||||
|
||||
def get_color_distortion(s=1.0):
|
||||
# s is the strength of color distortion.
|
||||
color_jitter = transforms.ColorJitter(0.8*s, 0.8*s, 0.8*s, 0.2*s)
|
||||
color_jitter = transforms.ColorJitter(0.8 * s, 0.8 * s, 0.8 * s, 0.2 * s)
|
||||
rnd_color_jitter = transforms.RandomApply([color_jitter], p=0.8)
|
||||
rnd_gray = transforms.RandomGrayscale(p=0.2)
|
||||
color_distort = transforms.Compose([
|
||||
rnd_color_jitter,
|
||||
rnd_gray])
|
||||
color_distort = transforms.Compose([rnd_color_jitter, rnd_gray])
|
||||
return color_distort
|
||||
|
||||
transform_list = []
|
||||
@@ -53,8 +50,25 @@ def make_transforms(
|
||||
return transform
|
||||
|
||||
|
||||
def make_transform_eval(
|
||||
crop_size=224, normalization=((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))
|
||||
):
|
||||
logger.info("making imagenet evaluation transforms")
|
||||
|
||||
resize_size = int(crop_size * 256 / 224)
|
||||
transform = transforms.Compose(
|
||||
[
|
||||
transforms.Resize(resize_size),
|
||||
transforms.CenterCrop(crop_size),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize(normalization[0], normalization[1]),
|
||||
]
|
||||
)
|
||||
return transform
|
||||
|
||||
|
||||
class GaussianBlur(object):
|
||||
def __init__(self, p=0.5, radius_min=0.1, radius_max=2.):
|
||||
def __init__(self, p=0.5, radius_min=0.1, radius_max=2.0):
|
||||
self.prob = p
|
||||
self.radius_min = radius_min
|
||||
self.radius_max = radius_max
|
||||
|
||||
Reference in New Issue
Block a user