Spaces:
Sleeping
Sleeping
| import os | |
| import copy | |
| import numpy as np | |
| from PIL import Image | |
| import random | |
| os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True" | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from torch.utils.data import Dataset | |
| from torchvision import transforms | |
| import torchvision.transforms.functional as TF | |
| import torchvision.models as tvm | |
| # ========================================== | |
| # CONFIG | |
| # ========================================== | |
| NUM_CLASSES = 7 | |
| IMAGENET_MEAN = [0.485, 0.456, 0.406] | |
| IMAGENET_STD = [0.229, 0.224, 0.225] | |
| COLOR_MAP = { | |
| (255, 255, 255): 0, # Non-change | |
| (0, 128, 0): 1, # Low Vegetation | |
| (0, 255, 0): 2, # Tree | |
| (128, 128, 128): 3, # N.V.G. Surface | |
| (0, 0, 255): 4, # Water | |
| (255, 0, 0): 5, # Playground | |
| (128, 0, 0): 6, # Building | |
| } | |
| DISPLAY_COLORS = { | |
| 0: (255, 255, 0), 1: (0, 100, 0), 2: (0, 255, 0), 3: (128, 128, 128), | |
| 4: (0, 0, 255), 5: (255, 0, 0), 6: (128, 0, 0) | |
| } | |
| LEGEND_ENTRIES = [(DISPLAY_COLORS[k], lbl) for k, lbl in enumerate( | |
| ["No Change", "Low Vegetation", "Trees", "N.V.G. Surface", | |
| "Water", "Playground", "Building"])] | |
| def rgb_to_class(mask): | |
| arr = np.array(mask) | |
| h, w, _ = arr.shape | |
| flat = arr.reshape(-1, 3) | |
| out = np.zeros(flat.shape[0], dtype=np.int64) | |
| for rgb, cls in COLOR_MAP.items(): | |
| out[np.all(flat == rgb, axis=1)] = cls | |
| return torch.from_numpy(out.reshape(h, w)).long() | |
| def class_map_to_rgb(cls_map): | |
| rgb = np.zeros((*cls_map.shape, 3), dtype=np.uint8) | |
| for cls, col in DISPLAY_COLORS.items(): | |
| rgb[cls_map == cls] = col | |
| return rgb | |
| def denormalize(tensor): | |
| mean = np.array(IMAGENET_MEAN, dtype=np.float32) | |
| std = np.array(IMAGENET_STD, dtype=np.float32) | |
| img = tensor.cpu().permute(1, 2, 0).numpy() | |
| img = (img * std + mean).clip(0, 1) | |
| return (img * 255).astype(np.uint8) | |
| # ========================================== | |
| # PAIRED AUGMENTATION | |
| # ========================================== | |
| class PairedAugment: | |
| def __init__(self, img_size=256): | |
| self.img_size = img_size | |
| self.jitter = transforms.ColorJitter( | |
| brightness=0.25, contrast=0.25, saturation=0.15, hue=0.04 | |
| ) | |
| self.blur = transforms.GaussianBlur(kernel_size=5, sigma=(0.1, 2.0)) | |
| def __call__(self, img1, img2, lbl1, lbl2): | |
| i, j, h, w = transforms.RandomResizedCrop.get_params( | |
| img1, scale=(0.5, 1.0), ratio=(1.0, 1.0) | |
| ) | |
| size = (self.img_size, self.img_size) | |
| img1 = TF.resized_crop(img1, i, j, h, w, size) | |
| img2 = TF.resized_crop(img2, i, j, h, w, size) | |
| lbl1 = TF.resized_crop(lbl1, i, j, h, w, size, | |
| interpolation=TF.InterpolationMode.NEAREST) | |
| lbl2 = TF.resized_crop(lbl2, i, j, h, w, size, | |
| interpolation=TF.InterpolationMode.NEAREST) | |
| if random.random() > 0.5: | |
| img1 = TF.hflip(img1); img2 = TF.hflip(img2) | |
| lbl1 = TF.hflip(lbl1); lbl2 = TF.hflip(lbl2) | |
| if random.random() > 0.7: | |
| img1 = TF.vflip(img1); img2 = TF.vflip(img2) | |
| lbl1 = TF.vflip(lbl1); lbl2 = TF.vflip(lbl2) | |
| if random.random() > 0.7: | |
| angle = random.choice([90, 180, 270]) | |
| img1 = TF.rotate(img1, angle); img2 = TF.rotate(img2, angle) | |
| lbl1 = TF.rotate(lbl1, angle); lbl2 = TF.rotate(lbl2, angle) | |
| img1 = self.jitter(img1) | |
| img2 = self.jitter(img2) | |
| if random.random() < 0.15: | |
| img1 = self.blur(img1) | |
| img2 = self.blur(img2) | |
| if random.random() < 0.05: | |
| img1 = TF.rgb_to_grayscale(img1, num_output_channels=3) | |
| img2 = TF.rgb_to_grayscale(img2, num_output_channels=3) | |
| return img1, img2, lbl1, lbl2 | |
| # ========================================== | |
| # DATASET | |
| # ========================================== | |
| class SECONDDataset(Dataset): | |
| def __init__(self, root_dir, transform=None, augment=False, mosaic_prob=0.4): | |
| self.dir_im1 = os.path.join(root_dir, 'im1') | |
| self.dir_im2 = os.path.join(root_dir, 'im2') | |
| self.dir_label1 = os.path.join(root_dir, 'label1') | |
| self.dir_label2 = os.path.join(root_dir, 'label2') | |
| self.names = sorted(os.listdir(self.dir_im1)) | |
| self.transform = transform | |
| self.augment = augment | |
| self.mosaic_prob = mosaic_prob if augment else 0.0 | |
| self.paired_aug = PairedAugment() if augment else None | |
| def __len__(self): | |
| return len(self.names) | |
| def _load_sample(self, idx): | |
| name = self.names[idx] | |
| img1 = Image.open(os.path.join(self.dir_im1, name)).convert('RGB') | |
| img2 = Image.open(os.path.join(self.dir_im2, name)).convert('RGB') | |
| l1 = Image.open(os.path.join(self.dir_label1, name)).convert('RGB') | |
| l2 = Image.open(os.path.join(self.dir_label2, name)).convert('RGB') | |
| return img1, img2, l1, l2 | |
| def _mosaic(self, idx): | |
| indices = [idx] + random.sample(range(len(self.names)), 3) | |
| half = 128 | |
| full = 256 | |
| canvas1 = Image.new('RGB', (full, full)) | |
| canvas2 = Image.new('RGB', (full, full)) | |
| c_lbl1 = Image.new('RGB', (full, full)) | |
| c_lbl2 = Image.new('RGB', (full, full)) | |
| positions = [(0, 0), (half, 0), (0, half), (half, half)] | |
| for (px, py), i in zip(positions, indices): | |
| im1, im2, lb1, lb2 = self._load_sample(i) | |
| im1 = im1.resize((half, half)) | |
| im2 = im2.resize((half, half)) | |
| lb1 = lb1.resize((half, half), Image.NEAREST) | |
| lb2 = lb2.resize((half, half), Image.NEAREST) | |
| canvas1.paste(im1, (px, py)) | |
| canvas2.paste(im2, (px, py)) | |
| c_lbl1.paste(lb1, (px, py)) | |
| c_lbl2.paste(lb2, (px, py)) | |
| return canvas1, canvas2, c_lbl1, c_lbl2 | |
| def __getitem__(self, idx): | |
| if self.augment and random.random() < self.mosaic_prob: | |
| img1, img2, l1, l2 = self._mosaic(idx) | |
| if random.random() > 0.5: | |
| img1 = TF.hflip(img1); img2 = TF.hflip(img2) | |
| l1 = TF.hflip(l1); l2 = TF.hflip(l2) | |
| if random.random() > 0.7: | |
| img1 = TF.vflip(img1); img2 = TF.vflip(img2) | |
| l1 = TF.vflip(l1); l2 = TF.vflip(l2) | |
| else: | |
| img1, img2, l1, l2 = self._load_sample(idx) | |
| img1 = img1.resize((256, 256)) | |
| img2 = img2.resize((256, 256)) | |
| l1 = l1.resize((256, 256), Image.NEAREST) | |
| l2 = l2.resize((256, 256), Image.NEAREST) | |
| if self.augment: | |
| img1, img2, l1, l2 = self.paired_aug(img1, img2, l1, l2) | |
| img1 = self.transform(img1) | |
| img2 = self.transform(img2) | |
| l1 = rgb_to_class(l1) | |
| l2 = rgb_to_class(l2) | |
| sem_change = torch.where(l1 == l2, torch.zeros_like(l2), l2) | |
| return img1, img2, sem_change | |
| # ========================================== | |
| # CBAM | |
| # ========================================== | |
| class ChannelAttention(nn.Module): | |
| def __init__(self, channels, reduction=8): | |
| super().__init__() | |
| self.avg_pool = nn.AdaptiveAvgPool2d(1) | |
| self.max_pool = nn.AdaptiveMaxPool2d(1) | |
| self.fc = nn.Sequential( | |
| nn.Linear(channels, channels // reduction, bias=False), | |
| nn.ReLU(inplace=True), | |
| nn.Linear(channels // reduction, channels, bias=False), | |
| ) | |
| self.sig = nn.Sigmoid() | |
| def forward(self, x): | |
| avg = self.fc(self.avg_pool(x).squeeze(-1).squeeze(-1)) | |
| mx = self.fc(self.max_pool(x).squeeze(-1).squeeze(-1)) | |
| return x * self.sig(avg + mx).view(x.size(0), x.size(1), 1, 1) | |
| class SpatialAttention(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.conv = nn.Conv2d(2, 1, kernel_size=7, padding=3, bias=False) | |
| self.sig = nn.Sigmoid() | |
| def forward(self, x): | |
| avg = x.mean(dim=1, keepdim=True) | |
| mx, _ = x.max(dim=1, keepdim=True) | |
| return x * self.sig(self.conv(torch.cat([avg, mx], dim=1))) | |
| class CBAM(nn.Module): | |
| def __init__(self, channels, reduction=8): | |
| super().__init__() | |
| self.ca = ChannelAttention(channels, reduction) | |
| self.sa = SpatialAttention() | |
| def forward(self, x): | |
| return self.sa(self.ca(x)) | |
| # ========================================== | |
| # BACKBONE (PRETRAINED RESNET34) | |
| # ========================================== | |
| class Backbone(nn.Module): | |
| def __init__(self, hidden=96, pretrained=True): | |
| super().__init__() | |
| h, h2 = hidden, hidden // 2 | |
| G = 8 | |
| try: | |
| weights = tvm.ResNet34_Weights.IMAGENET1K_V1 if pretrained else None | |
| resnet = tvm.resnet34(weights=weights) | |
| except AttributeError: | |
| resnet = tvm.resnet34(pretrained=pretrained) | |
| self.stem = nn.Sequential(resnet.conv1, resnet.bn1, resnet.relu, resnet.maxpool) | |
| self.layer1 = resnet.layer1 # 64ch, 64×64 | |
| self.layer2 = resnet.layer2 # 128ch, 32×32 | |
| self.layer3 = resnet.layer3 # 256ch, 16×16 | |
| self.proj1 = nn.Sequential(nn.Conv2d(64, h2, 1), nn.GroupNorm(G, h2), nn.ReLU(inplace=True)) | |
| self.proj2 = nn.Sequential(nn.Conv2d(128, h, 1), nn.GroupNorm(G, h), nn.ReLU(inplace=True)) | |
| self.proj3 = nn.Sequential(nn.Conv2d(256, h, 1), nn.GroupNorm(G, h), nn.ReLU(inplace=True)) | |
| def freeze_bn(self): | |
| for m in (self.stem, self.layer1, self.layer2, self.layer3): | |
| for layer in m.modules(): | |
| if isinstance(layer, nn.BatchNorm2d): | |
| layer.eval() | |
| def forward(self, x): | |
| x = self.stem(x) | |
| r1 = self.layer1(x) | |
| r2 = self.layer2(r1) | |
| r3 = self.layer3(r2) | |
| return self.proj1(r1), self.proj2(r2), self.proj3(r3) | |
| # ========================================== | |
| # CROSS-ATTENTION TEMPORAL FUSION | |
| # ========================================== | |
| class CrossAttentionTemporalFusion(nn.Module): | |
| def __init__(self, channels=96, num_heads=4, spatial_size=16, dropout=0.1): | |
| super().__init__() | |
| self.spatial_size = spatial_size | |
| self.cross_attn = nn.MultiheadAttention( | |
| channels, num_heads, batch_first=True, dropout=dropout | |
| ) | |
| self.diff_gate = nn.Sequential( | |
| nn.Conv2d(channels * 2, channels, 1), | |
| nn.GroupNorm(8, channels), | |
| nn.Sigmoid() | |
| ) | |
| self.norm1 = nn.LayerNorm(channels) | |
| self.norm2 = nn.LayerNorm(channels) | |
| self.ffn = nn.Sequential( | |
| nn.Linear(channels, channels * 2), | |
| nn.GELU(), | |
| nn.Dropout(dropout), | |
| nn.Linear(channels * 2, channels), | |
| ) | |
| self.norm3 = nn.LayerNorm(channels) | |
| def forward(self, f1, f2): | |
| f1 = F.adaptive_avg_pool2d(f1, (self.spatial_size, self.spatial_size)) | |
| f2 = F.adaptive_avg_pool2d(f2, (self.spatial_size, self.spatial_size)) | |
| B, C, H, W = f1.shape | |
| q = f2.flatten(2).permute(0, 2, 1) | |
| kv = f1.flatten(2).permute(0, 2, 1) | |
| q_norm = self.norm1(q) | |
| kv_norm = self.norm1(kv) | |
| attended, _ = self.cross_attn(q_norm, kv_norm, kv_norm) | |
| q = q + attended | |
| q = q + self.ffn(self.norm2(q)) | |
| out = self.norm3(q).permute(0, 2, 1).view(B, C, H, W) | |
| gate = self.diff_gate(torch.cat([f2, f1], dim=1)) | |
| return out * gate + f2 * (1.0 - gate) | |
| # ========================================== | |
| # DECODER | |
| # ========================================== | |
| class Decoder(nn.Module): | |
| def __init__(self, hidden=96, dropout=0.3): | |
| super().__init__() | |
| h, h2 = hidden, hidden // 2 | |
| G = 8 | |
| self.reduce = nn.Sequential( | |
| nn.Conv2d(3*h, 2*h, 3, 1, 1), nn.GroupNorm(G, 2*h), nn.ReLU(inplace=True), | |
| nn.Conv2d(2*h, h, 3, 1, 1), nn.GroupNorm(G, h), nn.ReLU(inplace=True), | |
| ) | |
| self.up1 = nn.Sequential( | |
| nn.Conv2d(h + h, h, 3, 1, 1), nn.GroupNorm(G, h), nn.ReLU(inplace=True), | |
| ) | |
| self.up2 = nn.Sequential( | |
| nn.Conv2d(h + h2, h, 3, 1, 1), nn.GroupNorm(G, h), nn.ReLU(inplace=True), | |
| ) | |
| self.up3 = nn.Sequential( | |
| nn.Conv2d(h, h2, 3, 1, 1), nn.GroupNorm(G, h2), nn.ReLU(inplace=True), | |
| ) | |
| self.up4 = nn.Sequential( | |
| nn.Conv2d(h2, h2, 3, 1, 1), nn.GroupNorm(G, h2), nn.ReLU(inplace=True), | |
| ) | |
| self.drop = nn.Dropout2d(p=dropout) | |
| self.sem_out = nn.Conv2d(h2, NUM_CLASSES, 1) | |
| self.bin_out = nn.Sequential( | |
| nn.Conv2d(h2, 32, 3, 1, 1), nn.GroupNorm(8, 32), nn.ReLU(inplace=True), | |
| nn.Dropout2d(p=0.2), nn.Conv2d(32, 2, 1), | |
| ) | |
| self.aux_head_32 = nn.Sequential( | |
| nn.Conv2d(h, h2, 3, 1, 1), nn.GroupNorm(G, h2), nn.ReLU(inplace=True), | |
| nn.Dropout2d(p=0.2), nn.Conv2d(h2, NUM_CLASSES, 1) | |
| ) | |
| self.aux_head_64 = nn.Sequential( | |
| nn.Conv2d(h, h2, 3, 1, 1), nn.GroupNorm(G, h2), nn.ReLU(inplace=True), | |
| nn.Dropout2d(p=0.2), nn.Conv2d(h2, NUM_CLASSES, 1) | |
| ) | |
| self.skip2_cbam = CBAM(h) | |
| self.skip1_cbam = CBAM(h2) | |
| def _gated_skip(self, cbam_module, feat_a, feat_b): | |
| diff = cbam_module(torch.abs(feat_a - feat_b)) | |
| gate = torch.sigmoid(diff) | |
| appear = (feat_a + feat_b) * 0.5 | |
| return appear * gate + diff | |
| def forward(self, x, s1_a, s1_b, s2_a, s2_b): | |
| skip2 = self._gated_skip(self.skip2_cbam, s2_a, s2_b) | |
| skip1 = self._gated_skip(self.skip1_cbam, s1_a, s1_b) | |
| x = self.reduce(x) | |
| x = F.interpolate(x, scale_factor=2, mode='bilinear', align_corners=False) | |
| x = torch.cat([x, F.adaptive_avg_pool2d(skip2, x.shape[2:])], dim=1) | |
| x = self.up1(x) | |
| aux32 = F.interpolate(self.aux_head_32(x), (256, 256), mode='bilinear', align_corners=False) | |
| x = F.interpolate(x, scale_factor=2, mode='bilinear', align_corners=False) | |
| x = torch.cat([x, F.adaptive_avg_pool2d(skip1, x.shape[2:])], dim=1) | |
| x = self.up2(x) | |
| aux64 = F.interpolate(self.aux_head_64(x), (256, 256), mode='bilinear', align_corners=False) | |
| x = F.interpolate(x, scale_factor=2, mode='bilinear', align_corners=False) | |
| x = self.up3(x) | |
| x = F.interpolate(x, scale_factor=2, mode='bilinear', align_corners=False) | |
| x = self.up4(x) | |
| x = self.drop(x) | |
| return self.sem_out(x), self.bin_out(x), aux32, aux64 | |
| # ========================================== | |
| # LOVÁSZ-SOFTMAX LOSS | |
| # ========================================== | |
| def lovasz_grad(gt_sorted): | |
| p = len(gt_sorted) | |
| gts = gt_sorted.sum() | |
| intersection = gts - gt_sorted.float().cumsum(0) | |
| union = gts + (1 - gt_sorted).float().cumsum(0) | |
| jaccard = 1.0 - intersection / union | |
| if p > 1: | |
| jaccard[1:p] = jaccard[1:p] - jaccard[0:-1] | |
| return jaccard | |
| def lovasz_softmax_flat(probs, labels, only_present=True): | |
| C = probs.shape[1] | |
| losses = [] | |
| for c in range(C): | |
| fg = (labels == c).float() | |
| if only_present and fg.sum() == 0: | |
| continue | |
| errors = (fg - probs[:, c]).abs() | |
| errors_s, perm = torch.sort(errors, 0, descending=True) | |
| fg_s = fg[perm] | |
| losses.append(torch.dot(errors_s, lovasz_grad(fg_s.detach()))) | |
| return torch.stack(losses).mean() if losses else probs.sum() * 0.0 | |
| class LovaszCELoss(nn.Module): | |
| def __init__(self, ce_weight=0.4, label_smoothing=0.02): | |
| super().__init__() | |
| self.ce_w = ce_weight | |
| self.label_smoothing = label_smoothing | |
| self.class_weights = None | |
| def forward(self, pred, target): | |
| probs = torch.softmax(pred, dim=1) | |
| B, C, H, W = probs.shape | |
| probs_flat = probs.permute(0, 2, 3, 1).reshape(-1, C) | |
| target_flat = target.reshape(-1) | |
| lov_loss = lovasz_softmax_flat(probs_flat, target_flat) | |
| ce_loss = F.cross_entropy( | |
| pred, target, | |
| weight=self.class_weights, | |
| label_smoothing=self.label_smoothing, | |
| ) | |
| return self.ce_w * ce_loss + (1.0 - self.ce_w) * lov_loss | |
| # ========================================== | |
| # DYNAMIC CLASS WEIGHTS | |
| # ========================================== | |
| class DynamicClassWeights: | |
| def __init__(self, num_classes, base_weights, device, ema_decay=0.95): | |
| self.base = torch.tensor(base_weights, dtype=torch.float32, device=device) | |
| self.f1_ema = torch.ones(num_classes, dtype=torch.float32, device=device) | |
| self.decay = ema_decay | |
| self.device = device | |
| def update(self, cm): | |
| eps = 1e-7 | |
| for c in range(len(self.f1_ema)): | |
| tp = cm[c, c] | |
| fp = cm[:, c].sum() - tp | |
| fn = cm[c, :].sum() - tp | |
| pr = tp / (tp + fp + eps) | |
| rc = tp / (tp + fn + eps) | |
| f1 = 2 * pr * rc / (pr + rc + eps) | |
| self.f1_ema[c] = self.decay * self.f1_ema[c] + (1.0 - self.decay) * float(f1) | |
| def weights(self): | |
| inv_f1 = 1.0 / (self.f1_ema + 0.1) | |
| scaled = inv_f1 / inv_f1.mean() | |
| return (self.base * scaled).clamp(0.5, 15.0) | |
| # ========================================== | |
| # MAIN MODEL (LSNN) | |
| # ========================================== | |
| class LSNN(nn.Module): | |
| def __init__(self, hidden=96): | |
| super().__init__() | |
| # Use the local checkpoint weights instead of downloading pretrained weights at startup. | |
| self.backbone = Backbone(hidden=hidden, pretrained=False) | |
| self.ltf = CrossAttentionTemporalFusion( | |
| channels=hidden, num_heads=4, spatial_size=16 | |
| ) | |
| self.cbam = CBAM(hidden) | |
| self.decoder = Decoder(hidden=hidden, dropout=0.3) | |
| def train(self, mode=True): | |
| super().train(mode) | |
| if mode: | |
| self.backbone.freeze_bn() | |
| return self | |
| def forward(self, i1, i2): | |
| s1_a, s2_a, s3_a = self.backbone(i1) | |
| s1_b, s2_b, s3_b = self.backbone(i2) | |
| ltc_out = self.cbam(self.ltf(s3_a, s3_b)) | |
| sp = ltc_out.shape[2:] | |
| f = torch.cat([ | |
| F.adaptive_avg_pool2d(s3_a, sp), | |
| F.adaptive_avg_pool2d(s3_b, sp), | |
| ltc_out, | |
| ], dim=1) | |
| return self.decoder(f, s1_a, s1_b, s2_a, s2_b) | |
| # ========================================== | |
| # EMA / SWA UTILITIES | |
| # ========================================== | |
| class ModelEMA: | |
| def __init__(self, model, decay=0.999): | |
| self.ema = copy.deepcopy(model).eval() | |
| self.decay = decay | |
| for p in self.ema.parameters(): | |
| p.requires_grad_(False) | |
| def hard_reset(self, model): | |
| for ema_p, p in zip(self.ema.parameters(), model.parameters()): | |
| ema_p.data.copy_(p.data) | |
| for ema_b, b in zip(self.ema.buffers(), model.buffers()): | |
| ema_b.copy_(b) | |
| def update(self, model, epoch=0, total_epochs=150): | |
| t = min(epoch / max(total_epochs - 1, 1), 1.0) | |
| decay = 0.990 + (0.9995 - 0.990) * t | |
| self.decay = decay | |
| for ema_p, p in zip(self.ema.parameters(), model.parameters()): | |
| ema_p.data.mul_(decay).add_(p.data, alpha=1.0 - decay) | |
| for ema_b, b in zip(self.ema.buffers(), model.buffers()): | |
| ema_b.copy_(b) | |
| class SWA: | |
| def __init__(self, model): | |
| self.avg = copy.deepcopy(model).eval() | |
| self.n = 0 | |
| for p in self.avg.parameters(): | |
| p.requires_grad_(False) | |
| def update(self, model): | |
| self.n += 1 | |
| for avg_p, p in zip(self.avg.parameters(), model.parameters()): | |
| avg_p.data.mul_(self.n / (self.n + 1)).add_(p.data / (self.n + 1)) | |
| for avg_b, b in zip(self.avg.buffers(), model.buffers()): | |
| avg_b.copy_(b) | |
| def reset(self, model): | |
| for avg_p, p in zip(self.avg.parameters(), model.parameters()): | |
| avg_p.data.copy_(p.data) | |
| for avg_b, b in zip(self.avg.buffers(), model.buffers()): | |
| avg_b.copy_(b) | |
| self.n = 1 | |
| # ========================================== | |
| # TEST-TIME AUGMENTATION (INFERENCE) | |
| # ========================================== | |
| def predict_tta(model, img1, img2): | |
| model.eval() | |
| with torch.no_grad(): | |
| s0, _, _, _ = model(img1, img2) | |
| # horizontal flip | |
| s1, _, _, _ = model(torch.flip(img1, [3]), torch.flip(img2, [3])) | |
| s1 = torch.flip(s1, [3]) | |
| # vertical flip | |
| s2, _, _, _ = model(torch.flip(img1, [2]), torch.flip(img2, [2])) | |
| s2 = torch.flip(s2, [2]) | |
| # 180° | |
| s3, _, _, _ = model(torch.flip(img1, [2, 3]), torch.flip(img2, [2, 3])) | |
| s3 = torch.flip(s3, [2, 3]) | |
| return (s0 + s1 + s2 + s3) * 0.25 |