{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":51753,"databundleVersionId":5692552,"sourceType":"competition"},{"sourceId":6360712,"sourceType":"datasetVersion","datasetId":3614787},{"sourceId":7821278,"sourceType":"datasetVersion","datasetId":4580103}],"dockerImageVersionId":30636,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"! pip install -q --no-index --find-links=file:///kaggle/input/contrails-base-model/pip \\\n    segmentation_models_pytorch\n! cp -r /kaggle/input/contrails-base-model/src ./src","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:19:28.463362Z","iopub.execute_input":"2024-03-25T15:19:28.46401Z","iopub.status.idle":"2024-03-25T15:19:43.75173Z","shell.execute_reply.started":"2024-03-25T15:19:28.463978Z","shell.execute_reply":"2024-03-25T15:19:43.750352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import Tuple\n\n\ndef get_range(ra: str, n: int) -> Tuple[int, int]:\n    \"\"\"\n    istart, iend = get_range('1:10', n)\n\n    Args:\n      ra (str or None): 1, 1:, 5:10\n      n (int): default iend when it is omitted, ra = '1:'\n    \"\"\"\n\n    if ra is None:\n        return 0, n\n    elif ra.isnumeric():\n        # Single number\n        i = int(ra)\n        return i, i + 1\n    elif ':' in ra:\n        v = ra.split(':')\n        assert len(v) == 2\n        istart = 0 if v[0] == '' else int(v[0])\n        iend = n if v[1] == '' else int(v[1])\n        return istart, iend\n    else:\n        raise ValueError('Failed to parse range: {}'.format(ra))\n\n\ndef as_list(s, *, dtype=int):\n    \"\"\"\n    Parse string of numbers to list of numbers\n\n    Arg:\n      s (str): '1,2,4...7' => [1, 2, 4, 5, 6, 7]\n    \"\"\"\n    ret = []\n    if not isinstance(s, str):\n        return [dtype(s), ]\n\n    for seg in s.split(','):\n        if dtype is int and '...' in seg:\n            start, last = map(int, seg.split('...'))\n            for i in range(start, last):\n                ret.append(dtype(i))\n        else:\n            ret.append(dtype(seg))\n\n    return ret","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:19:43.753804Z","iopub.execute_input":"2024-03-25T15:19:43.754169Z","iopub.status.idle":"2024-03-25T15:19:43.764636Z","shell.execute_reply.started":"2024-03-25T15:19:43.754134Z","shell.execute_reply":"2024-03-25T15:19:43.763786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport sys\nimport glob\nimport time\nimport yaml\nfrom tqdm.auto import tqdm\n# from sklearn.model_selection import KFold\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.nn.modules.loss import _Loss\n\nimport torchvision.transforms as T\nimport albumentations as A\nimport segmentation_models_pytorch as smp\n\nsys.path.append('/kaggle/working/src')\n# import util\nfrom lr_scheduler import Scheduler\n\n\ndi = '/kaggle/input/google-research-identify-contrails-reduce-global-warming'\ndevice = torch.device('cuda')","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:19:43.765927Z","iopub.execute_input":"2024-03-25T15:19:43.766542Z","iopub.status.idle":"2024-03-25T15:19:51.357475Z","shell.execute_reply.started":"2024-03-25T15:19:43.766517Z","shell.execute_reply":"2024-03-25T15:19:51.356511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cfg = yaml.safe_load(\"\"\"\ndata:\n  resize: 512          # maxvit_tiny_tf_512\n  augment: rotation    \n  augment_prob: 0.95    \n\nmodel:\n  encoder: tu-maxvit_tiny_tf_512.in1k   # timm-resnest26d\n  pretrained: False    \n  decoder_channels: [256, 128, 64, 32, 16]\n\nkfold:\n  k: 10\n  folds: 0  # 0,1,2,3,4\n\ntrain:\n  batch_size: 4        \n  weight_decay: 1e-2\n  clip_grad_norm: 1000.0\n  num_workers: 2\n\nval:\n  per_epoch: 1        \n  th: 0.45         \n\ntest:\n  th: 0.45           \n\nscheduler:\n  - linear:\n      lr_start: 1e-8\n      lr_end: 8e-4\n      epoch_end: 0.5\n  - cosine:\n      lr_end: 1e-6\n      epoch_end: 15   \n\"\"\")","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:19:51.359848Z","iopub.execute_input":"2024-03-25T15:19:51.36015Z","iopub.status.idle":"2024-03-25T15:19:51.370131Z","shell.execute_reply.started":"2024-03-25T15:19:51.360125Z","shell.execute_reply":"2024-03-25T15:19:51.369246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Regular grid for sampling y_sym from y\n# See torch grid_sample() for the convention\ndef create_grid(nc: int, offset=0.5) -> torch.Tensor:\n    \"\"\"\n    Create xy values of nc x nc grid\n    \n    Arg:\n      nc (int): number of grid points per dimension (e.g. 512)\n      offset (float): offset in sampling in units of original 256 x 256 image\n                      offset 0 gives identity mapping\n                      Use offset 0.5 for shifted contrail label\n\n    Returns: grid (Tensor)\n      grid points in [-1, 1] for torch grid_sample()\n    \"\"\"\n    grid = np.zeros((nc, nc, 2), dtype=np.float32)\n    for ix in range(nc):\n        for iy in range(nc):\n            grid[ix, iy, 1] = -1 + 2 * (ix + 0.5) / nc + offset / 128\n            grid[ix, iy, 0] = -1 + 2 * (iy + 0.5) / nc + offset / 128\n    grid = torch.from_numpy(grid).unsqueeze(0)\n    return grid\n\n\n# Augmentation\ndef augmentation(aug: str):\n    if aug == 'd4':  # Dihedral group D4\n        return A.Compose([\n            A.RandomRotate90(p=1),\n            A.HorizontalFlip(p=0.5),\n        ])\n    elif aug == 'rotation':\n        return A.Compose([\n            A.RandomRotate90(p=1),\n            A.HorizontalFlip(p=0.5),\n            A.ShiftScaleRotate(rotate_limit=30, scale_limit=0.2, p=0.75)\n        ])\n    else:\n        raise ValueError\n\n        \n# Create list of data\ndef create_df(data_type: str):\n    assert data_type in ['train', 'validation', 'test']\n\n    filenames = glob.glob('%s/%s/*' % (di, data_type))\n    filenames.sort()\n    assert filenames\n\n    file_ids = []\n    for filename in filenames:\n        file_id = os.path.basename(filename).replace('.h5', '')\n        file_ids.append(file_id)\n\n    return pd.DataFrame({'file_id': file_ids, 'filename': filenames})\n\n\n# Load from file\ndef load_data(path: str, t=4) -> dict:\n    \"\"\"\n    Arg:\n      path (str): directory including data npy\n\n    Returns: dict\n      x (array): (3, 256, 256) for bands 11, 14, 15\n      y (array): (1, 256, 256) for ground truth label\n      annotation_mean (Optional[array])\n    \"\"\"\n    file_id = path.split('/')[-1]\n\n    # Load input images\n    bands = []\n    for k in [11, 14, 15]:\n        a = np.load('%s/band_%02d.npy' % (path, k))\n        bands.append(a)\n\n    x = np.stack(bands, axis=0)  # (C, H, W, T)\n    \n    ret = {'file_id': file_id,\n           'x': x[:, :, :, t].copy()}\n\n    # Ground truth label\n    filename = '%s/human_pixel_masks.npy' % path\n    if os.path.exists(filename):\n        y = np.load(filename)\n        y_sum = np.sum(y)\n        ret['label'] = y.reshape(1, 256, 256).astype(np.float32)\n\n    # Soft label: mean of individual annotations\n    filename = '%s/human_individual_masks.npy' % path\n    if os.path.exists(filename):\n        annot = np.load(filename)\n        annot = annot.astype(np.float32)  # (256, 256, 1, A)\n        annot = np.mean(annot, axis=3).reshape(1, 256, 256)\n        \n        ret['annotation_mean'] = annot\n            \n    return ret\n\n\n# Ash false color\ndef rescale_range(x, f_min, f_max):\n    # Rescale [f_min, f_max] to [0, 1]\n    return (x - f_min) / (f_max - f_min)\n\n\ndef ash_color(x):\n    \"\"\"\n    False color for contrail annotation\n    x (array): (3, H, W) -> (3, H, W)\n    \"\"\"\n    r = rescale_range(x[2] - x[1], -4, 2)  # band 15 - 11\n    g = rescale_range(x[1] - x[0], -4, 5)  # band 14 - 11\n    b = rescale_range(x[1], 243, 303)      # band 14\n\n    x = torch.stack([r, g, b], axis=0)\n    x = 1 - x\n\n    return x\n\n\nclass Dataset(torch.utils.data.Dataset):\n    def __init__(self, df, cfg, *, augment=False):\n        self.df = df\n        \n        # Augmentation\n        self.augment = None\n        if augment and cfg['data']['augment']:\n            self.augment = augmentation(cfg['data']['augment'])\n\n        # Resize input image\n        nc = cfg['data']['resize']\n        self.resize = nn.Identity() if nc == 256 else T.Resize(nc, antialias=False)\n\n        # Sample y_sym from y on this 0.5-pixel shifted grid\n        self.grid = create_grid(nc, offset=0.5)\n\n        # The probability to apply augmentation (0.95)\n        self.augment_prob = cfg['data']['augment_prob']\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, i):\n        r = self.df.iloc[i]\n        file_id = r['file_id']\n        \n        d = load_data(r['filename'])\n\n        x = torch.from_numpy(d['x'])\n        y = label = None\n        if 'annotation_mean' in d:\n            y = torch.from_numpy(d['annotation_mean'])\n        \n        if 'label' in d:\n            label = torch.from_numpy(d['label'])\n\n        # Create false-color image\n        x = ash_color(x)\n        x = self.resize(x)\n\n        # Create symmetric label y_sym\n        if y is not None:\n            y_sym = F.grid_sample(y.unsqueeze(0), self.grid,\n                                  mode='bilinear', padding_mode='border',\n                                  align_corners=False).squeeze(0)\n\n        # Augment\n        w_original = 1.0\n        if self.augment is not None and np.random.random() < self.augment_prob:\n            assert y is not None  # y and y_sym always exist if you want to augment\n            w_original = 0.0      # augmented\n\n            x = x.permute(1, 2, 0).numpy()  # => (H, W, C)\n            y_sym = y_sym.permute(1, 2, 0).numpy()\n\n            aug = self.augment(image=x, mask=y_sym)\n\n            x = torch.from_numpy(aug['image'].transpose(2, 0, 1))     # => (C, H, W)\n            y_sym = torch.from_numpy(aug['mask'].transpose(2, 0, 1))  # array (1, 256, 256)\n\n        # Return values\n        ret = {'file_id': file_id,\n               'x': x,\n               'w': np.float32(w_original)} \n\n        # y is target for loss (soft label)\n        if y is not None:\n            ret['y_sym'] = y_sym\n            ret['y'] = y\n\n        # label is ground truth for score\n        if label is not None:\n            ret['label'] = label\n        return ret\n\n\nclass Data:\n    def __init__(self, data_type, *, debug=False):\n        # Load filename list\n        df = create_df(data_type)\n        if debug:\n            df = df.iloc[:100]\n\n        self.df = df\n\n    def __len__(self):\n        return len(self.df)\n\n    def dataset(self, cfg, augment):\n#         df = self.df.iloc[idx] if idx is not None else self.df\n\n        return Dataset(self.df, cfg, augment=augment)\n\n    def loader(self, cfg, *, augment=False, batch_size=None, shuffle=False, drop_last=False):\n        batch_size = batch_size if batch_size is not None else cfg['train']['batch_size']\n        num_workers = cfg['train']['num_workers']\n\n        ds = self.dataset( cfg, augment)\n        return torch.utils.data.DataLoader(ds,\n                                           batch_size=batch_size,\n                                           num_workers=num_workers,\n                                           shuffle=shuffle,\n                                           drop_last=drop_last)","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:19:51.371583Z","iopub.execute_input":"2024-03-25T15:19:51.372042Z","iopub.status.idle":"2024-03-25T15:19:51.405906Z","shell.execute_reply.started":"2024-03-25T15:19:51.372009Z","shell.execute_reply":"2024-03-25T15:19:51.405004Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n_tta_config = {'d4prob': (8, True), 'd4logit': (8, False),\n               'rotprob': (4, True), 'rotlogit': (4, False), 'none': (1, False)}\n\ndef _tta_stack(x: torch.Tensor, n: int) -> torch.Tensor:\n    \n    stack = []\n    for k in range(4):\n        xa = torch.rot90(x, k, dims=[2, 3])\n        stack.append(xa)\n        if n == 8:\n            stack.append(torch.flip(xa, dims=[3, ]))\n    return torch.cat(stack, dim=0)\n\n\ndef _tta_average(y_pred: torch.Tensor, n: int, prob: bool) -> torch.Tensor:\n    \n    batch_size, nch, H, W = y_pred.shape\n    y_pred = y_pred.view(n, batch_size // n, nch, H, W)\n    batch_size = batch_size // n\n    y_avg = torch.zeros((batch_size, 1, H, W), dtype=torch.float32, device=y_pred.device)\n\n    if prob:\n        y_pred = y_pred.sigmoid()\n\n    if n == 4:\n        for k in range(4):\n            y_avg += (1 / n) * torch.rot90(y_pred[k], -k, dims=[2, 3])\n    if n == 8:\n        for k in range(4):\n            y_avg += (1 / n) * torch.rot90(y_pred[2 * k], -k, dims=[2, 3])\n            y_avg += (1 / n) * torch.rot90(y_pred[2 * k + 1].flip(dims=[3, ]), -k, dims=(2, 3))\n\n    if prob:\n        y_avg = y_avg.clamp(1e-6, 1 - 1e-6)  # logit() is NaN at 0 and 1\n        return y_avg.logit()\n    return y_avg\n\n\nclass TTA:\n    \"\"\"\n    Test-Time Augment input and average output\n    \n    Args:\n      tta_str: d4prob, d4logit, rotprob, rotlogit\n    \"\"\"\n    def __init__(self, tta_str: str):        \n        self.n, self.prob = _tta_config[tta_str]\n    \n    def __repr__(self):\n        return 'TTA(n={}, prob={})'.format(self.n, self.prob)\n\n    def stack(self, x):\n        if self.n == 1:\n            return x\n        else:\n            return _tta_stack(x, self.n)\n\n    def average(self, y):\n        if self.n == 1:\n            return y\n        else:\n            return _tta_average(y, self.n, self.prob)","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:19:51.407072Z","iopub.execute_input":"2024-03-25T15:19:51.407421Z","iopub.status.idle":"2024-03-25T15:19:51.422369Z","shell.execute_reply.started":"2024-03-25T15:19:51.407389Z","shell.execute_reply":"2024-03-25T15:19:51.421489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef get_asym_conv(nc):\n    \"\"\"\n    Final tiny convolution from y_sym_pred to y_pred\n    - expected to shift 0.5 pixel back\n    - also reduce from 512 to 256 when y_sym is 512x512\n    \"\"\"\n    if nc == 256:\n        hidden_size = 9\n        asym_conv = nn.Sequential(\n            nn.Conv2d(1, hidden_size, kernel_size=(3, 3), padding=1, padding_mode='replicate'),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(hidden_size, 1, kernel_size=1),\n        )\n    elif nc == 512:\n        hidden_size = 25\n        asym_conv = nn.Sequential(\n            nn.Conv2d(1, hidden_size, kernel_size=(5, 5), padding=2, stride=2),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(hidden_size, 1, kernel_size=1),\n        )\n    else:\n        raise NotImplementedError\n\n    return asym_conv\n\n\nclass Model(nn.Module):\n    def __init__(self, cfg, pretrained=True, tta=None):\n        super().__init__()\n        name = cfg['model']['encoder']\n        pretrained = 'imagenet' if (pretrained and cfg['model']['pretrained']) else None\n        decoder_channels = cfg['model']['decoder_channels']  # (256, 128, 64, 32, 16)\n\n        self.unet = smp.Unet(name,\n                             encoder_weights=pretrained,\n                             classes=1,\n                             decoder_channels=decoder_channels,\n        )\n\n        self.asym_conv = get_asym_conv(cfg['data']['resize'])        \n        self.tta = TTA(tta) if tta is not None else None\n\n    def forward(self, x):\n#         if self.tta is not None:\n#             x = self.tta.stack(x)  # TTA input\n\n        y_sym = self.unet(x)\n\n#         if self.tta is not None:\n#             y_sym = self.tta.average(y_sym)\n            \n        # Tiny conv from y_sym_pred -> y_pred\n        y_pred = self.asym_conv(y_sym)\n\n        return y_sym, y_pred","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:19:51.424305Z","iopub.execute_input":"2024-03-25T15:19:51.424568Z","iopub.status.idle":"2024-03-25T15:19:51.43777Z","shell.execute_reply.started":"2024-03-25T15:19:51.424547Z","shell.execute_reply":"2024-03-25T15:19:51.43697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass BCELoss(_Loss):\n    def __init__(self):\n        super().__init__()\n        self.criterion = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self,\n                y_sym_pred: torch.Tensor, y_sym: torch.Tensor,\n                y_pred: torch.Tensor, y: torch.Tensor, w) -> torch.Tensor:\n        loss_sym = self.criterion(y_sym_pred, y_sym).mean(dim=(1, 2, 3))\n        loss_original = self.criterion(y_pred, y).mean(dim=(1, 2, 3))\n        loss = loss_sym + w * loss_original\n\n        return loss.mean()  # mean of batch\n\n\ndef evaluate(model, loader_val, *, th=0.45):\n    \"\"\"\n    Compute validation loss and score\n    \"\"\"\n    tb = time.time()\n\n    was_training = model.training\n    model.eval()\n\n    n_sum = 0\n    loss_sum = 0.0\n    dice_sum = 0.0\n    tp = 0\n    positives_pred = 0\n    positives_true = 0\n\n    for d in loader_val:\n        x = d['x'].to(device)      # input image (3, H, W); may be upscaled\n        if 'y' in d:\n            y = d['y'].to(device)  # soft label: (1, H, W)\n            y_sym = d['y_sym'].to(device)\n        else:\n            y = None\n\n        label = d['label'].to(device)  # ground-truth segmentation mask (always 256 x 256)\n        batch_size = len(x)\n\n        # Predict\n        with torch.no_grad():\n            y_sym_pred, y_pred = model(x)  # (batch_size, 1, H, W)\n\n        # Compute validation loss\n        if y is not None:\n            # No augmentation in validation but rescale criteion(y_pred, y) like training time\n            w = 1 - augment_fraction\n            loss = criterion(y_sym_pred, y_sym, y_pred, y, w)\n            loss_dice = dice(y_pred, y)\n \n            n_sum += batch_size\n            loss_sum += loss.item() * batch_size\n            dice_sum += loss_dice.item() * batch_size\n\n        # Compute score\n        y_pred = y_pred.sigmoid() > th\n        tp += (y_pred * label).sum().item()\n        positives_pred += y_pred.sum().item()\n        positives_true += label.sum().item()\n\n    global_dice = 2 * tp / (positives_pred + positives_true)\n\n    model.train(was_training)\n\n    dt = time.time() - tb\n    ret = {'score': global_dice,\n           'dt': dt}\n\n    if y is not None:\n        ret['loss'] = loss_sum / n_sum\n        ret['dice'] = 1 - dice_sum / n_sum\n    return ret","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:19:51.438756Z","iopub.execute_input":"2024-03-25T15:19:51.439007Z","iopub.status.idle":"2024-03-25T15:19:51.4546Z","shell.execute_reply.started":"2024-03-25T15:19:51.438986Z","shell.execute_reply":"2024-03-25T15:19:51.453744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"debug = True \n\n# Data\ndata = Data('train', debug=debug)\ndata_test = Data('validation', debug=debug)\nprint('Data', len(data), len(data_test))\n\n# Loss\ncriterion = BCELoss()\ndice = smp.losses.DiceLoss('binary', from_logits=True)  # just as metric in validation\n\nweight_decay = float(cfg['train']['weight_decay'])\nclip_grad_norm = float(cfg['train']['clip_grad_norm'])\n\naugment_fraction = cfg['data']['augment_prob']\nsteps_per_epoch = cfg['val']['per_epoch']\nth_val = cfg['val']['th']\nth_test = cfg['test']['th']\n\n# Train loop\nlog = {}\nepochs_log = []\nlosses_train = []\nfinal_scores = []\n\nn_sum = 0\nloss_sum = 0.0\nlrs = []\n\nprint('Pretrained: ', cfg['model']['pretrained'])\n\ntb_global = time.time()\n\nloader_train = data.loader( cfg, augment=False, drop_last=True, shuffle=True)\n# loader_val = data.loader(idx_val, cfg)\nloader_test = data_test.loader(cfg)\n\nnbatch = len(loader_train)\n\n# Model\nmodel = Model(cfg, pretrained=True)\nmodel.to(device)\nmodel.train()\n\n# Optimizer\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-4,\n                              weight_decay=weight_decay)\n\nscheduler = Scheduler(optimizer, cfg['scheduler'])\nepochs = 4 if debug else len(scheduler)\nprint('%d epochs' % epochs)","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:19:51.455734Z","iopub.execute_input":"2024-03-25T15:19:51.455988Z","iopub.status.idle":"2024-03-25T15:19:54.270968Z","shell.execute_reply.started":"2024-03-25T15:19:51.455967Z","shell.execute_reply":"2024-03-25T15:19:54.270039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tb = time.time()\ndt_val = 0.0\nprint('       Epoch loss  dice score lr time')\nfor iepoch in range(epochs):\n    # This is for validating multiple time per epoch\n    icheck = [nbatch * (i + 1) // steps_per_epoch - 1 for i in range(steps_per_epoch)]\n\n    for ibatch, d in enumerate(loader_train):\n        x = d['x'].to(device)          # input image\n        y = d['y'].to(device)          # segmentation label\n        y_sym = d['y_sym'].to(device)  # symmetric label\n        w = d['w'].to(device)          # w=0 if augmented, 1 if x and y are original \n        batch_size = len(x)\n\n        optimizer.zero_grad()\n\n        # Predict\n        y_sym_pred, y_pred = model(x)  # (batch_size, 1, 256, 256)\n        loss = criterion(y_sym_pred, y_sym, y_pred, y, w)\n\n        # Backpropagate\n        loss.backward()\n\n        n_sum += batch_size\n        loss_sum += batch_size * loss.item()\n\n        nn.utils.clip_grad_norm_(model.parameters(), clip_grad_norm)\n        optimizer.step()\n\n        ep = iepoch + (ibatch + 1) / nbatch\n        lr = optimizer.param_groups[0]['lr']\n        lrs.append((ep, lr))\n\n        # Validation\n        if ibatch == icheck[0]:\n            icheck.pop(0)\n\n            epochs_log.append(iepoch + (ibatch + 1) / nbatch)\n            loss_train = loss_sum / n_sum\n            losses_train.append(loss_train)\n\n#             val = evaluate(model, loader_val, th=th_val)\n            test = evaluate(model, loader_test, th=th_test)\n\n#             losses_val.append(val['loss'])\n            dt = time.time() - tb\n            dt_val += test['dt']\n\n            print('Epoch %5.2f %6.3f %.3f  %5.1e %5.1f %5.1f min' % (ep,\n                  10 * loss_train\n                  , test['score'],\n                  lr, dt_val / 60, dt / 60))\n\n            # Reset train loss\n            n_sum = 0\n            loss_sum = 0.0\n\n        scheduler.step(ep)\n\n# Save model\nmodel.eval()\nofilename = '/kaggle/working/model_7_3_2024_1.pytorch' \ntorch.save(model.state_dict(), ofilename)\nprint(ofilename, 'written')\ndel model\n\n# Kfolds done\ndt = time.time() - tb_global\nprint('Total time: %.2f min' % (dt / 60))","metadata":{"execution":{"iopub.status.busy":"2024-03-25T15:19:54.274006Z","iopub.execute_input":"2024-03-25T15:19:54.274349Z","iopub.status.idle":"2024-03-25T15:29:55.763874Z","shell.execute_reply.started":"2024-03-25T15:19:54.274321Z","shell.execute_reply":"2024-03-25T15:29:55.762444Z"},"trusted":true},"execution_count":null,"outputs":[]}]}