{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":39763,"databundleVersionId":11756775,"sourceType":"competition"}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# This was the code used to generate the simple further finetuned models from bartley's original work\n\n# The original work was https://www.kaggle.com/code/brendanartley/convnext-full-resolution-baseline?scriptVersionId=241259000\n\n# The models of this work are publically available at https://www.kaggle.com/datasets/harshitsheoran/simple-further-finetuned-bartley-open-models\n\n# The code below was run locally and is not an implementation to run on kaggle notebooks, please take ideas from the code and implement in your own pipelines instead of trying to run it directly as that might not work without many modifications.","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!/usr/bin/env python\n# coding: utf-8\n\n# In[1]:\n\n\ndef is_dist_avail_and_initialized():\n    if not dist.is_available():\n        return False\n    if not dist.is_initialized():\n        return False\n    return True\n\n\ndef get_world_size():\n    if not is_dist_avail_and_initialized():\n        return 1\n    return dist.get_world_size()\n\n\ndef get_rank():\n    if not is_dist_avail_and_initialized():\n        return 0\n    return dist.get_rank()\n\n\ndef is_main_process():\n    return get_rank() == 0\n\n\ndef save_on_master(*args, **kwargs):\n    if is_main_process():\n        torch.save(*args, **kwargs)\n\n\ndef setup_for_distributed(is_master):\n    \"\"\"\n    This function disables printing when not in master process\n    \"\"\"\n    import builtins as __builtin__\n    builtin_print = __builtin__.print\n\n    def print(*args, **kwargs):\n        force = kwargs.pop('force', False)\n        if is_master or force:\n            builtin_print(*args, **kwargs)\n\n    __builtin__.print = print\n    \ndef init_distributed():\n\n    # Initializes the distributed backend which will take care of sychronizing nodes/GPUs\n    dist_url = \"env://\" # default\n    # only works with torch.distributed.launch // torch.run\n    rank = int(os.environ[\"RANK\"])\n    world_size = int(os.environ['WORLD_SIZE'])\n    local_rank = int(os.environ['LOCAL_RANK'])\n    \n    #print('init process group')\n    dist.init_process_group(\n            backend=\"nccl\",\n            init_method=dist_url,\n            world_size=world_size,\n            rank=rank)\n    #print('done init process group')\n    \n    # this will make all .cuda() calls work properly\n    try:\n        torch.cuda.set_device(local_rank)\n    except:\n        print(\"error at\", local_rank)\n    # synchronizes all the threads to reach this point before moving on\n    #print('incoming barrier')\n    dist.barrier()\n    #print(\"got through it, into setup\")\n    setup_for_distributed(rank == 0)\n    #print(\"done setup\")\n    \ndef seed_everything(seed=1234):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.use_deterministic_algorithms(True)\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nfrom tqdm import tqdm\nimport os\nfrom glob import glob\nimport copy\nimport time\nimport math\nimport command\nimport random\nimport sys\nimport h5py\n\nos.environ['CUBLAS_WORKSPACE_CONFIG'] = ':4096:8'\n\n#os.environ['CUDA_VISIBLE_DEVICES'] = '1,2'\nos.environ['NO_ALBUMENTATIONS_UPDATE'] = '1'\n\nimport cv2\nfrom PIL import Image\nimport matplotlib as mpl\nimport matplotlib.pyplot as plt\nmpl.rcParams['figure.figsize'] = 12, 8\n\nfrom skimage import img_as_ubyte\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom sklearn.model_selection import *\nfrom sklearn.metrics import *\n\nimport torch\nfrom torch import nn, optim\nfrom torch.utils.data import Dataset, DataLoader\nimport segmentation_models_pytorch as smp\nimport timm\nfrom torchtoolbox.tools import mixup_data, mixup_criterion\nimport torchvision as tv\nfrom torch_ema import ExponentialMovingAverage\n\nfrom transformers import get_cosine_schedule_with_warmup\n\nimport torch.distributed as dist\n\nimport logging\nlogging.getLogger('timm').setLevel(logging.WARNING)\n\nimport webdataset as wds\n\nimport torch.multiprocessing\ntorch.multiprocessing.set_sharing_strategy('file_system')\n\n\n# In[ ]:\n\n\n\n\n\n# In[125]:\n\n\nclass CFG:\n    world_size = get_world_size()\n    rank = get_rank()\n    \n    DDP = 1\n    DDP_INIT_DONE = 0\n    N_GPUS = 4\n    FOLD = 0\n    FULLDATA = 0\n    \n    COMPILE = True\n    \n    model_name = 'convnext_small.fb_in22k_ft_in1k'\n    V = '4'\n    \n    OUTPUT_FOLDER = f\"./data/AAA_SEG/TRY6_SEG/{model_name}_v{V}\"\n    \n    seed = 3407\n    \n    device = torch.device('cuda')\n    \n    n_folds = 4\n    \n    image_size = [1000, 70]\n    \n    train_batch_size = 32\n    valid_batch_size = 32\n    acc_steps = 1\n    \n    lr = 1e-4\n    wd = 1e-3\n    ema_decay_per_epoch = 0.3\n    freeze_epochs = 0\n    n_epochs = 40\n    n_cycles = 1\n    n_warmup_steps = 0\n    upscale_steps = 1.6\n    validate_every = 1\n    \n    epoch = 0\n    global_step = 0\n    literal_step = 0\n    \n    autocast = True\n    \n    workers = 0\n\nif CFG.FULLDATA:\n    CFG.seed = CFG.FOLD\n    \nOUTPUT_FOLDER = CFG.OUTPUT_FOLDER\n        \nCFG.cache_dir = CFG.OUTPUT_FOLDER + f'/cache/'\nos.makedirs(CFG.cache_dir, exist_ok=1)\n\nseed_everything(CFG.seed)\n\n\n# In[ ]:\n\n\n\n\n\n# In[126]:\n\n# My own dataset, it is as follows\ndata = pd.read_csv('./data/OpenFWI_sampled_velfp32.csv')\n#    seis_path\t                                            vel_path\t                                                    method      fn\t    isample\n#0\t./data/OpenFWI_sampled//CurveVel_A_data57_sample0.npy\t./data/OpenFWI_sampled_velfp32//CurveVel_A_model57_sample0.npy\tCurveVel_A\tdata57\t0\n#1\t./data/OpenFWI_sampled//CurveVel_A_data57_sample1.npy\t./data/OpenFWI_sampled_velfp32//CurveVel_A_model57_sample1.npy\tCurveVel_A\tdata57\t1\n#2\t./data/OpenFWI_sampled//CurveVel_A_data57_sample2.npy\t./data/OpenFWI_sampled_velfp32//CurveVel_A_model57_sample2.npy\tCurveVel_A\tdata57\t2\n#3\t./data/OpenFWI_sampled//CurveVel_A_data57_sample3.npy\t./data/OpenFWI_sampled_velfp32//CurveVel_A_model57_sample3.npy\tCurveVel_A\tdata57\t3\n#4\t./data/OpenFWI_sampled//CurveVel_A_data57_sample4.npy\t./data/OpenFWI_sampled_velfp32//CurveVel_A_model57_sample4.npy\tCurveVel_A\tdata57\t4\n\nbartley_folds = pd.read_csv('./data/bartley_folds.csv')\nfpath_to_fold = dict(zip(bartley_folds.data_fpath, bartley_folds.fold))\n\ndata['data_fpath'] = data['method'] + \"/\" + data['fn']+'.npy'\n\ndata['bartley_fold'] = data.data_fpath.apply(lambda x: fpath_to_fold[x])\n\ndata\n\n\n# In[ ]:\n\n\n\n\n\n# In[127]:\n\n\nmethods = sorted(data.method.unique())\nmethod_to_idx = {v: i for i, v in enumerate(methods)}\nmethod_to_idx\n\n\n# In[ ]:\n\n\n\n\n\n# In[128]:\n\n\nclass SeismicDataset(Dataset):\n    def __init__(self, data, transforms=None, is_training=False):\n        self.data = data\n        self.transforms = transforms\n        self.is_training = is_training\n        \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, i):\n        row = self.data.iloc[i]\n        \n        image = np.load(row.seis_path)\n        mask = np.load(row.vel_path)\n        \n        mask = torch.as_tensor(mask).float()\n        \n        image = torch.as_tensor(image).float()\n        \n        label = torch.zeros((10,))\n        \n        return {\n            'images': image,\n            'masks': mask,\n            'labels': label,\n            'ids': f\"{row.seis_path}\",\n            'methods': row.method\n        }\n\n\n# In[ ]:\n\n\n\n\n\n# In[129]:\n\ndef get_loaders(ret_data=False, n_workers=CFG.workers):\n    \n    train_data = data[data.bartley_fold!=0].reset_index(drop=True)\n    valid_data = data[data.bartley_fold==0].reset_index(drop=True)\n    \n    train_dataset = SeismicDataset(train_data, None, 1)\n    valid_dataset = SeismicDataset(valid_data, None, 0)\n    \n    if CFG.DDP and CFG.DDP_INIT_DONE:\n        train_sampler = torch.utils.data.distributed.DistributedSampler(dataset=train_dataset, shuffle=True, drop_last=True)\n        train_sampler.set_epoch(CFG.epoch) #needed for shuffling?\n        \n        train_loader = DataLoader(train_dataset, batch_size=CFG.train_batch_size, sampler=train_sampler, num_workers=CFG.workers, pin_memory=True, drop_last=True)\n        \n        valid_sampler = torch.utils.data.distributed.DistributedSampler(dataset=valid_dataset, shuffle=False)\n        valid_loader = DataLoader(valid_dataset, batch_size=CFG.valid_batch_size, sampler=valid_sampler, shuffle=False, num_workers=CFG.workers, pin_memory=True)\n    else:\n        train_loader = DataLoader(train_dataset, batch_size=CFG.train_batch_size, shuffle=True, num_workers=n_workers, pin_memory=False)\n        valid_loader = DataLoader(valid_dataset, batch_size=CFG.valid_batch_size, shuffle=False, num_workers=n_workers, pin_memory=False)\n    \n    CFG.steps_per_epoch = math.ceil(len(train_loader) / CFG.acc_steps)\n    \n    if ret_data:\n        return train_loader, valid_loader, train_data, valid_data\n    return train_loader, valid_loader\n    \ntrain_loader, valid_loader, train_data, valid_data = get_loaders(ret_data=True, n_workers=0)\n\nfor d in valid_loader: break\n\n_, axs = plt.subplots(1, 4, figsize=(30, 15))\naxs = axs.flatten()\nfor img, ax in zip(range(4), axs):\n    try:\n        ax.imshow(d['images'][img].numpy()[:3].transpose(1, 2, 0), vmin=-1.5, vmax=1.5, cmap='gray')\n    except: pass\n    \n_, axs = plt.subplots(1, 4, figsize=(30, 15))\naxs = axs.flatten()\nfor img, ax in zip(range(4), axs):\n    try:\n        ax.imshow(d['masks'][img].numpy()[:3].transpose(1, 2, 0), cmap='gray')\n    except: pass\n\n\n# In[ ]:\n\n\n\n\n\n# In[ ]:\n\n\n\n\n\n# In[ ]:\n\n\n\n\n\n# In[113]:\n\n\nfrom types import MethodType\nfrom timm.models.convnext import ConvNeXtBlock\nfrom monai.networks.blocks import UpSample, SubpixelUpsample\n\nclass ConvBnAct2d(nn.Module):\n    def __init__(\n        self,\n        in_channels,\n        out_channels,\n        kernel_size,\n        padding: int = 0,\n        stride: int = 1,\n        norm_layer: nn.Module = nn.Identity,\n        act_layer: nn.Module = nn.ReLU,\n    ):\n        super().__init__()\n\n        self.conv= nn.Conv2d(\n            in_channels, \n            out_channels,\n            kernel_size,\n            stride=stride, \n            padding=padding, \n            bias=False,\n        )\n        self.norm = norm_layer(out_channels) if norm_layer != nn.Identity else nn.Identity()\n        self.act= act_layer(inplace=True)\n\n    def forward(self, x):\n        x = self.conv(x)\n        x = self.norm(x)\n        x = self.act(x)\n        return x\n\n\nclass SCSEModule2d(nn.Module):\n    def __init__(self, in_channels, reduction=16):\n        super().__init__()\n        self.cSE = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Conv2d(in_channels, in_channels // reduction, 1),\n            nn.Tanh(),\n            nn.Conv2d(in_channels // reduction, in_channels, 1),\n            nn.Sigmoid(),\n        )\n        self.sSE = nn.Sequential(\n            nn.Conv2d(in_channels, 1, 1), \n            nn.Sigmoid(),\n            )\n\n    def forward(self, x):\n        return x * self.cSE(x) + x * self.sSE(x)\n\nclass Attention2d(nn.Module):\n    def __init__(self, name, **params):\n        super().__init__()\n        if name is None:\n            self.attention = nn.Identity(**params)\n        elif name == \"scse\":\n            self.attention = SCSEModule2d(**params)\n        else:\n            raise ValueError(\"Attention {} is not implemented\".format(name))\n\n    def forward(self, x):\n        return self.attention(x)\n\nclass DecoderBlock2d(nn.Module):\n    def __init__(\n        self,\n        in_channels,\n        skip_channels,\n        out_channels,\n        norm_layer: nn.Module = nn.Identity,\n        attention_type: str = None,\n        intermediate_conv: bool = False,\n        upsample_mode: str = \"deconv\",\n        scale_factor: int = 2,\n    ):\n        super().__init__()\n\n        # Upsample block\n        if upsample_mode == \"pixelshuffle\":\n            self.upsample= SubpixelUpsample(\n                spatial_dims= 2,\n                in_channels= in_channels,\n                scale_factor= scale_factor,\n            )\n        else:\n            self.upsample = UpSample(\n                spatial_dims= 2,\n                in_channels= in_channels,\n                out_channels= in_channels,\n                scale_factor= scale_factor,\n                mode= upsample_mode,\n            )\n\n        if intermediate_conv:\n            k= 3\n            c= skip_channels if skip_channels != 0 else in_channels\n            self.intermediate_conv = nn.Sequential(\n                ConvBnAct2d(c, c, k, k//2),\n                ConvBnAct2d(c, c, k, k//2),\n                )\n        else:\n            self.intermediate_conv= None\n\n        self.attention1 = Attention2d(\n            name= attention_type, \n            in_channels= in_channels + skip_channels,\n            )\n\n        self.conv1 = ConvBnAct2d(\n            in_channels + skip_channels,\n            out_channels,\n            kernel_size= 3,\n            padding= 1,\n            norm_layer= norm_layer,\n        )\n\n        self.conv2 = ConvBnAct2d(\n            out_channels,\n            out_channels,\n            kernel_size= 3,\n            padding= 1,\n            norm_layer= norm_layer,\n        )\n        self.attention2 = Attention2d(\n            name= attention_type, \n            in_channels= out_channels,\n            )\n\n    def forward(self, x, skip=None):\n        x = self.upsample(x)\n\n        if self.intermediate_conv is not None:\n            if skip is not None:\n                skip = self.intermediate_conv(skip)\n            else:\n                x = self.intermediate_conv(x)\n\n        if skip is not None:\n            #print(x.shape, skip.shape)\n            x = torch.cat([x, skip], dim=1)\n            x = self.attention1(x)\n\n        x = self.conv1(x)\n        x = self.conv2(x)\n        x = self.attention2(x)\n        return x\n\n\nclass UnetDecoder2d(nn.Module):\n    \"\"\"\n    Unet decoder.\n    Source: https://arxiv.org/abs/1505.04597\n    \"\"\"\n    def __init__(\n        self,\n        encoder_channels: tuple[int],\n        skip_channels: tuple[int] = None,\n        decoder_channels: tuple = (256, 128, 64, 32),\n        scale_factors: tuple = (2,2,2,2),\n        norm_layer: nn.Module = nn.Identity,\n        attention_type: str = None,\n        intermediate_conv: bool = False,\n        upsample_mode: str = \"deconv\",\n    ):\n        super().__init__()\n        \n        if len(encoder_channels) == 4:\n            decoder_channels= decoder_channels[1:]\n        self.decoder_channels= decoder_channels\n        \n        if skip_channels is None:\n            skip_channels= list(encoder_channels[1:]) + [0]\n\n        # Build decoder blocks\n        in_channels= [encoder_channels[0]] + list(decoder_channels[:-1])\n        self.blocks = nn.ModuleList()\n\n        for i, (ic, sc, dc) in enumerate(zip(in_channels, skip_channels, decoder_channels)):\n            # print(i, ic, sc, dc)\n            self.blocks.append(\n                DecoderBlock2d(\n                    ic, sc, dc, \n                    norm_layer= norm_layer,\n                    attention_type= attention_type,\n                    intermediate_conv= intermediate_conv,\n                    upsample_mode= upsample_mode,\n                    scale_factor= scale_factors[i],\n                    )\n            )\n\n    def forward(self, feats: list[torch.Tensor]):\n        res= [feats[0]]\n        feats= feats[1:]\n\n        # Decoder blocks\n        for i, b in enumerate(self.blocks):\n            skip= feats[i] if i < len(feats) else None\n            res.append(\n                b(res[-1], skip=skip),\n                )\n            \n        return res\n\nclass SegmentationHead2d(nn.Module):\n    def __init__(\n        self,\n        in_channels,\n        out_channels,\n        scale_factor: tuple[int] = (2,2),\n        kernel_size: int = 3,\n        mode: str = \"nontrainable\",\n    ):\n        super().__init__()\n        self.conv= nn.Conv2d(\n            in_channels, out_channels, kernel_size= kernel_size,\n            padding= kernel_size//2\n        )\n        self.upsample = UpSample(\n            spatial_dims= 2,\n            in_channels= out_channels,\n            out_channels= out_channels,\n            scale_factor= scale_factor,\n            mode= mode,\n        )\n\n    def forward(self, x):\n        x = self.conv(x)\n        x = self.upsample(x)\n        return x\n        \n\n#############\n## Encoder ##\n#############\n\ndef _convnext_block_forward(self, x):\n    shortcut = x\n    x = self.conv_dw(x)\n\n    if self.use_conv_mlp:\n        x = self.norm(x)\n        x = self.mlp(x)\n    else:\n        x = self.norm(x)\n        x = x.permute(0, 2, 3, 1)\n        x = x.contiguous()\n        x = self.mlp(x)\n        x = x.permute(0, 3, 1, 2)\n        x = x.contiguous()\n\n    if self.gamma is not None:\n        x = x * self.gamma.reshape(1, -1, 1, 1)\n\n    x = self.drop_path(x) + self.shortcut(shortcut)\n    return x\n\n\nclass Model(nn.Module):\n    def __init__(\n        self,\n        backbone: str = False,\n        pretrained: bool = True,\n    ):\n        if not backbone:\n            backbone = CFG.model_name\n        \n        super().__init__()\n        \n        # Encoder\n        self.backbone= timm.create_model(\n            backbone,\n            in_chans= 5,\n            pretrained= pretrained,\n            features_only= True,\n            drop_path_rate=0.0,\n            )\n        ecs= [_[\"num_chs\"] for _ in self.backbone.feature_info][::-1]\n\n        # Decoder\n        self.decoder= UnetDecoder2d(\n            encoder_channels= ecs,\n        )\n\n        self.seg_head= SegmentationHead2d(\n            in_channels= self.decoder.decoder_channels[-1],\n            out_channels= 1,\n            scale_factor= 1,\n        )\n        \n        self._update_stem(backbone)\n        \n        self.replace_activations(self.backbone, log=False)\n        self.replace_norms(self.backbone, log=False)\n        self.replace_forwards(self.backbone, log=False)\n\n    def _update_stem(self, backbone):\n        if backbone.startswith(\"convnext\"):\n\n            # Update stride\n            self.backbone.stem_0.stride = (4, 1)\n            self.backbone.stem_0.padding = (0, 2)\n\n            # Duplicate stem layer (to downsample height)\n            with torch.no_grad():\n                w = self.backbone.stem_0.weight\n                new_conv= nn.Conv2d(w.shape[0], w.shape[0], kernel_size=(4, 4), stride=(4, 1), padding=(0, 1))\n                new_conv.weight.copy_(w.repeat(1, (128//w.shape[1])+1, 1, 1)[:, :new_conv.weight.shape[1], :, :])\n                new_conv.bias.copy_(self.backbone.stem_0.bias)\n\n            self.backbone.stem_0= nn.Sequential(\n                nn.ReflectionPad2d((1,1,80,80)),\n                self.backbone.stem_0,\n                new_conv,\n            )\n\n        else:\n            raise ValueError(\"Custom striding not implemented.\")\n        pass\n\n    def replace_activations(self, module, log=False):\n        if log:\n            print(f\"Replacing all activations with GELU...\")\n        \n        # Apply activations\n        for name, child in module.named_children():\n            if isinstance(child, (\n                nn.ReLU, nn.LeakyReLU, nn.Mish, nn.Sigmoid, \n                nn.Tanh, nn.Softmax, nn.Hardtanh, nn.ELU, \n                nn.SELU, nn.PReLU, nn.CELU, nn.GELU, nn.SiLU,\n            )):\n                setattr(module, name, nn.GELU())\n            else:\n                self.replace_activations(child)\n\n    def replace_norms(self, mod, log=False):\n        if log:\n            print(f\"Replacing all norms with InstanceNorm...\")\n            \n        for name, c in mod.named_children():\n\n            # Get feature size\n            n_feats= None\n            if isinstance(c, (nn.BatchNorm2d, nn.InstanceNorm2d)):\n                n_feats= c.num_features\n            elif isinstance(c, (nn.GroupNorm,)):\n                n_feats= c.num_channels\n            elif isinstance(c, (nn.LayerNorm,)):\n                n_feats= c.normalized_shape[0]\n\n            if n_feats is not None:\n                new = nn.InstanceNorm2d(\n                    n_feats,\n                    affine=True,\n                    )\n                setattr(mod, name, new)\n            else:\n                self.replace_norms(c)\n\n    def replace_forwards(self, mod, log=False):\n        if log:\n            print(f\"Replacing forward functions...\")\n            \n        for name, c in mod.named_children():\n            if isinstance(c, ConvNeXtBlock):\n                c.forward = MethodType(_convnext_block_forward, c)\n            else:\n                self.replace_forwards(c)\n\n        \n    def proc_flip(self, x_in):\n        x_in= torch.flip(x_in, dims=[-3, -1])\n        x= self.backbone(x_in)\n        x= x[::-1]\n\n        # Decoder\n        x= self.decoder(x)\n        x_seg= self.seg_head(x[-1])\n        x_seg= x_seg[..., 1:-1, 1:-1]\n        x_seg= torch.flip(x_seg, dims=[-1])\n        x_seg= x_seg * 1500 + 3000\n        return x_seg\n\n    def forward(self, batch):\n        x= batch\n\n        # Encoder\n        x_in = x\n        x= self.backbone(x)\n        # print([_.shape for _ in x])\n        x= x[::-1]\n\n        # Decoder\n        x= self.decoder(x)\n        # print([_.shape for _ in x])\n        x_seg= self.seg_head(x[-1])\n        x_seg= x_seg[..., 1:-1, 1:-1]\n        x_seg= x_seg * 1500 + 3000\n        \n        x_seg = nn.functional.interpolate(x_seg, (70, 70), mode='bilinear')\n        \n        #return None, x_seg\n        \n        if self.training:\n            return None, x_seg\n        else:\n            p1 = self.proc_flip(x_in)\n            x_seg = torch.mean(torch.stack([x_seg, p1]), dim=0)\n            return None, x_seg\n\n\n# In[114]:\n\n\nif CFG.model_name==-1: CFG.model_name = 'convnext_small.fb_in22k_ft_in1k'\n\n#'''\nmodel = Model()\n\n_ = model.eval()\n\n#'''\ninp = d['images'][:2]\n\ninp = nn.functional.interpolate(inp, CFG.image_size, mode='bilinear').float()\n\nouts = model(inp)\n\n_ = [print(o.shape) for o in outs if o!=None]\n#'''\n\n\n# In[ ]:\n\n\n\n\n\n# In[98]:\n\n\nclass CustomLoss(nn.Module):\n    def __init__(self):\n        super(CustomLoss, self).__init__()\n        \n        self.mae = nn.L1Loss()\n        \n    def forward(self, outputs=None, targets=None, outputs_masks=None, targets_masks=None):\n        loss = 0.\n        \n        if targets!=None:\n            # This is meant to work if any aux is available (never got used in this comp)\n            loss1 = self.mae(outputs, targets)\n            loss = loss + (loss1 * 100)\n        \n        if targets_masks!=None:\n            #Only this one works\n            loss2 = self.mae(outputs_masks, targets_masks)\n            \n            loss = loss + (loss2 * 1)\n        \n        return loss\n\ndef plot_lr():\n    m = nn.Linear(2, 1)\n    optimizer = optim.AdamW(m.parameters(), lr=CFG.lr, weight_decay=CFG.wd)\n    scheduler = get_cosine_schedule_with_warmup(optimizer, num_training_steps=CFG.steps_per_epoch * CFG.n_epochs * CFG.upscale_steps, num_warmup_steps=CFG.n_warmup_steps)\n    \n    lrs = []\n    for s in range(int(CFG.n_epochs*CFG.steps_per_epoch*CFG.upscale_steps)):\n        lr = optimizer.param_groups[0]['lr']\n        scheduler.step()\n        lrs.append(lr)\n        if s==CFG.n_epochs*CFG.steps_per_epoch:\n            break\n    return lrs\n\ndef define_criterion_optimizer_scheduler_scaler(model):\n    criterion = CustomLoss().cuda()\n    \n    optimizer = optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.wd, fused=True)\n    \n    ema_decay_per_iter = CFG.ema_decay_per_epoch ** (1 / CFG.steps_per_epoch)\n    ema = ExponentialMovingAverage(model.parameters(), decay=ema_decay_per_iter)\n    \n    schedulers = [get_cosine_schedule_with_warmup(optimizer, num_training_steps=CFG.steps_per_epoch * CFG.n_epochs * CFG.upscale_steps, num_warmup_steps=CFG.n_warmup_steps) \n                  for _ in range(CFG.n_cycles)]\n    \n    scaler = torch.cuda.amp.GradScaler(enabled=CFG.autocast)\n    \n    return criterion, optimizer, schedulers, scaler, ema\n\n\n# In[ ]:\n\n\n\n# In[ ]:\n\n\n\n\n\n# In[100]:\n\n\nTRAIN_BATCH_CACHE_DICT, VALID_BATCH_CACHE_DICT = {}, {}\n\ndef train_one_epoch(model, loader, running_dist=True):\n    model.train()\n    running_loss = 0.0\n    last20loss = []\n    \n    if is_main_process(): bar = tqdm(loader, bar_format='{n_fmt}/{total_fmt} {elapsed}<{remaining} {postfix}', position=0, leave=True)\n    else: bar = loader\n    \n    batch_cache = False\n    if CFG.epoch!=0:\n        if is_main_process():\n            bar = tqdm(range(len(loader)), bar_format='{n_fmt}/{total_fmt} {elapsed}<{remaining} {postfix}', position=0, leave=True)\n        else:\n            bar = range(len(loader))\n            \n        batch_cache = True\n        \n        train_keys = list(TRAIN_BATCH_CACHE_DICT.keys())\n        np.random.shuffle(train_keys)\n    \n    for step, data in enumerate(bar):\n        #break\n        step += 1\n        \n        if len(last20loss)>20:\n            last20loss.pop(0)\n        \n        if not batch_cache:\n            images = data['images']\n            targets = None #data['labels'].cuda()\n            targets_masks = data['masks']\n            \n            images = nn.functional.interpolate(images, CFG.image_size, mode='bilinear')\n            \n            TRAIN_BATCH_CACHE_DICT[f\"{step}_train\"] = [images, targets, targets_masks]\n        else:\n            #images, targets, targets_masks = TRAIN_BATCH_CACHE_DICT[f\"{step}_train\"]\n            \n            images, targets, targets_masks = TRAIN_BATCH_CACHE_DICT[train_keys[step-1]]\n        \n        images = images.cuda(non_blocking=True)\n        targets_masks = targets_masks.cuda(non_blocking=True)\n        \n        #Here is where the horizontal flip happens\n        B = images.size(0)\n        flip_mask = torch.rand(B, device=images.device) < 0.5\n        images[flip_mask] = images[flip_mask].flip(dims=[1, -1])\n        targets_masks[flip_mask] = targets_masks[flip_mask].flip(dims=[-1])\n        \n        with torch.cuda.amp.autocast(enabled=CFG.autocast, dtype=torch.float16):\n            logits, logits_masks = model(images)\n            \n            loss = criterion(logits, targets, logits_masks, targets_masks)\n        \n        running_loss += (loss.item() - running_loss) * (1 / step)\n        last20loss.append(loss.item())\n        \n        loss = loss / CFG.acc_steps\n        scaler.scale(loss).backward()\n        \n        if step % CFG.acc_steps == 0 or step == len(bar):\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            scheduler.step()\n            ema.update()\n            CFG.global_step += 1\n        \n        CFG.literal_step += 1\n        \n        lr = \"{:2e}\".format(optimizer.param_groups[0]['lr'])\n        \n        if is_main_process():\n            bar.set_postfix(loss=running_loss, last20loss=np.mean(last20loss), lr=float(lr), step=CFG.global_step)\n        \n        if running_dist:\n            dist.barrier()\n        \n        if step==10: break\n    \n    if is_main_process():\n        if running_dist:\n            if CFG.COMPILE:\n                torch.save(model.module._orig_mod.state_dict(), f\"{OUTPUT_FOLDER}/{CFG.FOLD}.pth\")\n            else:\n                torch.save(model.module.state_dict(), f\"{OUTPUT_FOLDER}/{CFG.FOLD}.pth\")\n            \n        else:\n            if CFG.COMPILE:\n                torch.save(model._orig_mod.state_dict(), f\"{OUTPUT_FOLDER}/{CFG.FOLD}.pth\")\n            else:\n                torch.save(model.state_dict(), f\"{OUTPUT_FOLDER}/{CFG.FOLD}.pth\")\n        \n        \ndef valid_one_epoch(path, loader, running_dist=True, debug=False, do_ema=True):\n    model = Model(pretrained=False)\n    st = torch.load(path, map_location=f\"cpu\")\n    model.eval()\n    model.cuda()\n    model.load_state_dict(st, strict=False)\n    \n    if is_main_process(): bar = tqdm(loader, bar_format='{n_fmt}/{total_fmt} {elapsed}<{remaining} {postfix}', position=0, leave=True)\n    else: bar = loader\n    \n    running_loss = 0.\n    \n    OUTPUTS, TARGETS, IDS, scores = [], [], [], []\n    \n    method_to_scores = {method: [] for method in methods}\n    \n    batch_cache = False\n    if CFG.epoch!=0:\n        if is_main_process():\n            bar = tqdm(range(len(loader)), bar_format='{n_fmt}/{total_fmt} {elapsed}<{remaining} {postfix}', position=0, leave=True)\n        else:\n            bar = range(len(loader))\n            \n        batch_cache = True\n    \n    for step, data in enumerate(bar):\n        \n        with torch.no_grad():\n            \n            if not batch_cache:\n                images = data['images']\n                targets = None #data['labels'].cuda()\n                targets_masks = data['masks']\n                ids = data['ids']\n                methods_ = data['methods']\n                \n                images = nn.functional.interpolate(images, CFG.image_size, mode='bilinear')\n                \n                VALID_BATCH_CACHE_DICT[f\"{step}_valid\"] = [images, targets, targets_masks, ids, methods_]\n            else:\n                images, targets, targets_masks, ids, methods_ = VALID_BATCH_CACHE_DICT[f\"{step}_valid\"]\n            \n            images = images.cuda(non_blocking=True)\n            targets_masks = targets_masks.cuda(non_blocking=True)\n            \n            with torch.cuda.amp.autocast(enabled=CFG.autocast):\n                if do_ema:\n                    with ema.average_parameters():\n                        logits, logits_mask = model(images)\n                else:\n                    logits, logits_mask = model(images)\n            \n            outputs = logits_mask.float().detach().cpu().numpy()\n            targets = targets_masks.float().detach().cpu().numpy()\n            \n            #'''\n            if running_dist:\n                dist.barrier()\n                \n                np.save(f'{CFG.cache_dir}/preds_{get_rank()}.npy', outputs)\n                np.save(f'{CFG.cache_dir}/targets_{get_rank()}.npy', targets)a\n                np.save(f'{CFG.cache_dir}/ids_{get_rank()}.npy', ids)\n                \n                dist.barrier()\n                \n                if is_main_process():\n                    outputs = np.concatenate([np.load(f\"{CFG.cache_dir}/preds_{_}.npy\") for _ in range(CFG.N_GPUS)])\n                    targets = np.concatenate([np.load(f\"{CFG.cache_dir}/targets_{_}.npy\") for _ in range(CFG.N_GPUS)])\n                    ids = np.concatenate([np.load(f\"{CFG.cache_dir}/ids_{_}.npy\") for _ in range(CFG.N_GPUS)])\n                    \n                dist.barrier()\n            else:    \n                pass\n            \n            for target, output, method in zip(targets, outputs, methods_):\n                score = np.abs(target-output).mean()\n                \n                method_to_scores[method].append(score)\n                \n                scores.append(score)\n            \n            if step==10: break\n            \n    if running_dist:\n        dist.barrier()\n    \n    if is_main_process():\n        score = np.mean(scores)\n        \n        print(f\"EPOCH {CFG.epoch+1} | MAE {score}\")\n        for method in method_to_scores:\n            print(f\"{method}: {np.mean(method_to_scores[method])}\")\n        \n        if debug:\n            return score, OUTPUTS, TARGETS, IDS\n    \n        return score\n    \n    if debug:\n        return [], [], [], []\n    \ndef run(model, get_loaders):\n    if is_main_process():\n        epochs = []\n        scores = []\n    \n    best_score = float('inf')\n    \n    for cycle in range(1, CFG.n_cycles+1):\n        global scheduler\n        scheduler = schedulers[cycle-1]\n        \n        for epoch in range(int(CFG.n_epochs*(cycle-1)), int(CFG.n_epochs*cycle)):\n            CFG.epoch = epoch\n\n            train_loader, valid_loader = get_loaders()\n\n            train_one_epoch(model, train_loader, running_dist=CFG.DDP_INIT_DONE)\n\n            if CFG.DDP_INIT_DONE:\n                dist.barrier()\n\n            if (CFG.epoch+1)%CFG.validate_every==0 or epoch==0:\n                if is_main_process():\n                    with ema.average_parameters():\n                        if CFG.DDP_INIT_DONE:\n                            if CFG.COMPILE:\n                                torch.save(model.module._orig_mod.state_dict(), f\"{OUTPUT_FOLDER}/{CFG.FOLD}_EMA.pth\")\n                            else:\n                                torch.save(model.module.state_dict(), f\"{OUTPUT_FOLDER}/{CFG.FOLD}_EMA.pth\")\n                        else:\n                            if CFG.COMPILE:\n                                torch.save(model.module._orig_mod.state_dict(), f\"{OUTPUT_FOLDER}/{CFG.FOLD}_EMA.pth\")\n                            else:\n                                torch.save(model.state_dict(), f\"{OUTPUT_FOLDER}/{CFG.FOLD}_EMA.pth\")\n\n                if CFG.DDP_INIT_DONE:\n                    dist.barrier()\n\n                score = valid_one_epoch(f\"{OUTPUT_FOLDER}/{CFG.FOLD}_EMA.pth\", valid_loader, debug=False, running_dist=CFG.DDP_INIT_DONE)\n\n            if CFG.DDP_INIT_DONE:\n                dist.barrier()\n\n            if is_main_process():\n                epochs.append(epoch)\n                scores.append(score)\n\n                if score <= best_score:\n                    print(\"SAVING BEST!\")\n                    if CFG.DDP_INIT_DONE:\n                        with ema.average_parameters():\n                            if CFG.COMPILE:\n                                torch.save(model.module._orig_mod.state_dict(), f\"{OUTPUT_FOLDER}/{CFG.FOLD}_best.pth\")\n                            else:\n                                torch.save(model.module.state_dict(), f\"{OUTPUT_FOLDER}/{CFG.FOLD}_best.pth\")\n                    else:\n                        with ema.average_parameters():\n                            if CFG.COMPILE:\n                                torch.save(model.module._orig_mod.state_dict(), f\"{OUTPUT_FOLDER}/{CFG.FOLD}_best.pth\")\n                            else:\n                                torch.save(model.state_dict(), f\"{OUTPUT_FOLDER}/{CFG.FOLD}_best.pth\")\n\n                    best_score = score        \n\n                try:\n                    command.run(['rm', '-r', CFG.cache_dir])\n                    pass\n                except:\n                    pass\n\n                os.makedirs(CFG.cache_dir, exist_ok=1)\n\n\n# In[ ]:\n\n\n\n\n\n# In[ ]:\n\n\nCFG.DDP = 1\n\nif __name__ == '__main__' and CFG.DDP:\n    \n    world_size = init_distributed()\n    CFG.DDP_INIT_DONE = 1\n    \n    local_rank = int(os.environ['LOCAL_RANK'])\n    \n    CFG.world_size = get_world_size()\n    CFG.rank = local_rank\n    \n    #important to setup before defining scheduler to establish the correct number of steps per epoch\n    train_loader, valid_loader = get_loaders()\n    \n    model = Model().cuda()\n    \n    st = torch.load('./data/pretrained/bartley_unet2d_convnext_seed2_epochbest.pt', map_location='cpu')\n    new_st = {}\n    for key in st:\n        new_st[key.replace('_orig_mod.', '')] = st[key]\n    model.load_state_dict(new_st)\n    \n    model = torch.compile(model)\n    \n    model = nn.parallel.DistributedDataParallel(model, device_ids=[local_rank], find_unused_parameters=True)\n    \n    #\n    \n    criterion, optimizer, schedulers, scaler, ema = define_criterion_optimizer_scheduler_scaler(model)\n    \n    run(model, get_loaders)\n    \nelse:\n    print(\"Please Run in DDP\")\n    \nimport sys\nsys.exit(0)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}