{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport pandas as pd\nimport os\nimport numpy as np\nfrom torch.utils.data.sampler import WeightedRandomSampler\nfrom timm.scheduler import CosineLRScheduler\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom torch.cuda.amp import GradScaler, autocast\nfrom tqdm import tqdm\nimport copy\nimport time\nimport os\nimport timm\nimport random\nimport numpy as np\nimport gc\nimport torch\nimport torchaudio\nimport torchvision\nfrom sklearn.model_selection import StratifiedKFold\nfrom metrics import calculate_competition_metrics, metrics_to_string, calculate_competition_metrics_no_map\nfrom warmup_scheduler import GradualWarmupScheduler\nfrom torch.optim import AdamW\nimport albumentations as A\nimport matplotlib.pyplot as plt\nfrom pylab import rcParams\n\nimport torch\nfrom torch.autograd import Variable\nimport torch.nn.functional as F\nimport numpy as np\ntry:\n    from itertools import  ifilterfalse\nexcept ImportError: # py3k\n    from itertools import  filterfalse as ifilterfalse\n\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n","metadata":{"ExecuteTime":{"end_time":"2024-05-02T08:55:35.545387Z","start_time":"2024-05-02T08:55:32.908334Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training Logs","metadata":{}},{"cell_type":"markdown","source":"Epoch 1 - Train loss: 0.8522, Train cmAP_1 : 0.5357, Train cmAP_5 : 0.7442, \nEpoch 1 - Valid loss: 0.8485, Valid cmAP_1 : 0.1220, Valid cmAP_5 : 0.3257, Valid mAP : 0.0082, Valid ROC : 0.5372, \nEpoch 1 - Save Best Score: 0.5372 Model\n\nEpoch 2 - Train loss: 0.7673, Train cmAP_1 : 0.5354, Train cmAP_5 : 0.7419, \nEpoch 2 - Valid loss: 0.6464, Valid cmAP_1 : 0.1235, Valid cmAP_5 : 0.3270, Valid mAP : 0.0129, Valid ROC : 0.5611, \nEpoch 2 - Save Best Score: 0.5611 Model\n\nEpoch 3 - Train loss: 0.3128, Train cmAP_1 : 0.5624, Train cmAP_5 : 0.7448, \nEpoch 3 - Valid loss: 0.1051, Valid cmAP_1 : 0.0590, Valid cmAP_5 : 0.2511, Valid mAP : 0.0226, Valid ROC : 0.6110, \nEpoch 3 - Save Best Score: 0.6110 Model\n\nEpoch 4 - Train loss: 0.0411, Train cmAP_1 : 0.5409, Train cmAP_5 : 0.7417, \nEpoch 4 - Valid loss: 0.0401, Valid cmAP_1 : 0.1435, Valid cmAP_5 : 0.3504, Valid mAP : 0.0463, Valid ROC : 0.7263, \nEpoch 4 - Save Best Score: 0.7263 Model\n\nEpoch 5 - Train loss: 0.0350, Train cmAP_1 : 0.5427, Train cmAP_5 : 0.7459, \nEpoch 5 - Valid loss: 0.0328, Valid cmAP_1 : 0.2460, Valid cmAP_5 : 0.4553, Valid mAP : 0.2485, Valid ROC : 0.8743, \nEpoch 5 - Save Best Score: 0.8743 Model\n\nEpoch 6 - Train loss: 0.0288, Train cmAP_1 : 0.5864, Train cmAP_5 : 0.7732, \nEpoch 6 - Valid loss: 0.0261, Valid cmAP_1 : 0.3153, Valid cmAP_5 : 0.5258, Valid mAP : 0.3372, Valid ROC : 0.9202, \nEpoch 6 - Save Best Score: 0.9202 Model\n\nEpoch 7 - Train loss: 0.0256, Train cmAP_1 : 0.5644, Train cmAP_5 : 0.7617, \nEpoch 7 - Valid loss: 0.0274, Valid cmAP_1 : 0.3493, Valid cmAP_5 : 0.5602, Valid mAP : 0.3776, Valid ROC : 0.9380, \nEpoch 7 - Save Best Score: 0.9380 Model\n\nEpoch 8 - Train loss: 0.0236, Train cmAP_1 : 0.5792, Train cmAP_5 : 0.7725, \nEpoch 8 - Valid loss: 0.0198, Valid cmAP_1 : 0.4676, Valid cmAP_5 : 0.6471, Valid mAP : 0.5536, Valid ROC : 0.9509, \nEpoch 8 - Save Best Score: 0.9509 Model\n\nEpoch 9 - Train loss: 0.0223, Train cmAP_1 : 0.6052, Train cmAP_5 : 0.7826, \nEpoch 9 - Valid loss: 0.0202, Valid cmAP_1 : 0.4754, Valid cmAP_5 : 0.6594, Valid mAP : 0.5519, Valid ROC : 0.9569, \nEpoch 9 - Save Best Score: 0.9569 Model\n\nEpoch 10 - Train loss: 0.0217, Train cmAP_1 : 0.5950, Train cmAP_5 : 0.7810, \nEpoch 10 - Valid loss: 0.0198, Valid cmAP_1 : 0.4619, Valid cmAP_5 : 0.6518, Valid mAP : 0.5328, Valid ROC : 0.9592, \nEpoch 10 - Save Best Score: 0.9592 Model\n\nEpoch 11 - Train loss: 0.0202, Train cmAP_1 : 0.6104, Train cmAP_5 : 0.7915, \nEpoch 11 - Valid loss: 0.0166, Valid cmAP_1 : 0.5612, Valid cmAP_5 : 0.7150, Valid mAP : 0.6656, Valid ROC : 0.9641, \nEpoch 11 - Save Best Score: 0.9641 Model\n\nEpoch 12 - Train loss: 0.0197, Train cmAP_1 : 0.5942, Train cmAP_5 : 0.7841, \nEpoch 12 - Valid loss: 0.0162, Valid cmAP_1 : 0.5966, Valid cmAP_5 : 0.7349, Valid mAP : 0.6805, Valid ROC : 0.9634, \nValid loss didn't improve last 1 epochs.\n\nEpoch 13 - Train loss: 0.0192, Train cmAP_1 : 0.5937, Train cmAP_5 : 0.7795, \nEpoch 13 - Valid loss: 0.0159, Valid cmAP_1 : 0.6153, Valid cmAP_5 : 0.7457, Valid mAP : 0.6893, Valid ROC : 0.9669, \nEpoch 13 - Save Best Score: 0.9669 Model\n\nEpoch 14 - Train loss: 0.0191, Train cmAP_1 : 0.5991, Train cmAP_5 : 0.7903, \nEpoch 14 - Valid loss: 0.0193, Valid cmAP_1 : 0.5310, Valid cmAP_5 : 0.7047, Valid mAP : 0.5732, Valid ROC : 0.9660, \nValid loss didn't improve last 1 epochs.\n\nEpoch 15 - Train loss: 0.0180, Train cmAP_1 : 0.6038, Train cmAP_5 : 0.7822, \nEpoch 15 - Valid loss: 0.0163, Valid cmAP_1 : 0.5753, Valid cmAP_5 : 0.7283, Valid mAP : 0.6613, Valid ROC : 0.9656, \nValid loss didn't improve last 2 epochs.\n\nEpoch 16 - Train loss: 0.0181, Train cmAP_1 : 0.6067, Train cmAP_5 : 0.7918, \nEpoch 16 - Valid loss: 0.0168, Valid cmAP_1 : 0.5555, Valid cmAP_5 : 0.7140, Valid mAP : 0.6448, Valid ROC : 0.9613, \nValid loss didn't improve last 3 epochs.\n\nEpoch 17 - Train loss: 0.0176, Train cmAP_1 : 0.6181, Train cmAP_5 : 0.7918, \nEpoch 17 - Valid loss: 0.0160, Valid cmAP_1 : 0.5999, Valid cmAP_5 : 0.7436, Valid mAP : 0.6738, Valid ROC : 0.9670, \nEpoch 17 - Save Best Score: 0.9670 Model\n\nEpoch 18 - Train loss: 0.0169, Train cmAP_1 : 0.6040, Train cmAP_5 : 0.7897, \nEpoch 18 - Valid loss: 0.0174, Valid cmAP_1 : 0.5920, Valid cmAP_5 : 0.7447, Valid mAP : 0.6454, Valid ROC : 0.9690, \nEpoch 18 - Save Best Score: 0.9690 Model\n\nEpoch 19 - Train loss: 0.0162, Train cmAP_1 : 0.6005, Train cmAP_5 : 0.7870, \nEpoch 19 - Valid loss: 0.0174, Valid cmAP_1 : 0.5855, Valid cmAP_5 : 0.7397, Valid mAP : 0.6386, Valid ROC : 0.9682, \nValid loss didn't improve last 1 epochs.\n\nEpoch 20 - Train loss: 0.0161, Train cmAP_1 : 0.6355, Train cmAP_5 : 0.8080, \nEpoch 20 - Valid loss: 0.0151, Valid cmAP_1 : 0.6357, Valid cmAP_5 : 0.7629, Valid mAP : 0.7092, Valid ROC : 0.9683, \nValid loss didn't improve last 2 epochs.\n\nEpoch 21 - Train loss: 0.0157, Train cmAP_1 : 0.6019, Train cmAP_5 : 0.7856, \nEpoch 21 - Valid loss: 0.0162, Valid cmAP_1 : 0.5980, Valid cmAP_5 : 0.7464, Valid mAP : 0.6594, Valid ROC : 0.9675, \nValid loss didn't improve last 3 epochs.\n\nEpoch 22 - Train loss: 0.0160, Train cmAP_1 : 0.5569, Train cmAP_5 : 0.7717, \nEpoch 22 - Valid loss: 0.0148, Valid cmAP_1 : 0.6402, Valid cmAP_5 : 0.7668, Valid mAP : 0.7204, Valid ROC : 0.9673, \nValid loss didn't improve last 4 epochs.\n\nEpoch 23 - Train loss: 0.0158, Train cmAP_1 : 0.6211, Train cmAP_5 : 0.8011, \nEpoch 23 - Valid loss: 0.0150, Valid cmAP_1 : 0.6254, Valid cmAP_5 : 0.7594, Valid mAP : 0.7124, Valid ROC : 0.9642, \nValid loss didn't improve last 5 epochs.\n\nEpoch 24 - Train loss: 0.0153, Train cmAP_1 : 0.5914, Train cmAP_5 : 0.7817, \nEpoch 24 - Valid loss: 0.0154, Valid cmAP_1 : 0.6207, Valid cmAP_5 : 0.7590, Valid mAP : 0.6824, Valid ROC : 0.9650, \nValid loss didn't improve last 6 epochs.\n\nEpoch 25 - Train loss: 0.0143, Train cmAP_1 : 0.6056, Train cmAP_5 : 0.7859, \nEpoch 25 - Valid loss: 0.0155, Valid cmAP_1 : 0.6094, Valid cmAP_5 : 0.7542, Valid mAP : 0.6769, Valid ROC : 0.9680, \nValid loss didn't improve last 7 epochs.\n\nEarly stop, Training End.\n","metadata":{}},{"cell_type":"code","source":"\nexp_name = 'exp1\nbackbone = 'eca_nfnet_l0'\nseed = 42\nbatch_size = 64\nnum_workers = 0\n\nn_epochs = 100\nwarmup_epo = 5\ncosine_epo = n_epochs - warmup_epo\n\nimage_size = 256\n\nlr_max = 1e-5\nlr_min = 1e-7\nweight_decay = 1e-6\n\nmel_spec_params = {\n    \"sample_rate\": 32000,\n    \"n_mels\": 128,\n    \"f_min\": 20,\n    \"f_max\": 16000,\n    \"n_fft\": 2048,\n    \"hop_length\": 512,\n    \"normalized\": True,\n    \"center\" : True,\n    \"pad_mode\" : \"constant\",\n    \"norm\" : \"slaney\",\n    \"onesided\" : True,\n    \"mel_scale\" : \"slaney\"\n}\n\ntop_db = 80\ntrain_period = 5\nval_period = 5\n\nsecondary_coef = 1.0\n\ntrain_duration = train_period * mel_spec_params[\"sample_rate\"]\nval_duration = val_period * mel_spec_params[\"sample_rate\"]\n\nN_FOLD = 5\nfold = 2\n\nuse_amp = True\nmax_grad_norm = 10\nearly_stopping = 7\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\noutput_folder = \"outputs\"\nos.makedirs(output_folder, exist_ok=True)\nos.makedirs(os.path.join(output_folder, exp_name), exist_ok=True)\n\n","metadata":{"ExecuteTime":{"end_time":"2024-05-02T08:55:35.550548Z","start_time":"2024-05-02T08:55:35.546391Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Seed Everything","metadata":{}},{"cell_type":"code","source":"def set_seed(seed=42):\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.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    \nset_seed(seed)\n","metadata":{"ExecuteTime":{"end_time":"2024-05-02T08:55:35.567584Z","start_time":"2024-05-02T08:55:35.550548Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv('../data/birdclef-2024/train_metadata.csv')\ndf[\"path\"] = \"../data/birdclef-2024/train_audio/\" + df[\"filename\"]\ndf[\"rating\"] = np.clip(df[\"rating\"] / df[\"rating\"].max(), 0.1, 1.0)\n\nskf = StratifiedKFold(n_splits=N_FOLD, random_state=seed, shuffle=True)\ndf['fold'] = -1\nfor ifold, (train_idx, val_idx) in enumerate(skf.split(X=df, y=df[\"primary_label\"].values)):\n    df.loc[val_idx, 'fold'] = ifold\n\nsub = pd.read_csv(\"../data/birdclef-2024/sample_submission.csv\")\ntarget_columns = sub.columns.tolist()[1:]\nnum_classes = len(target_columns)\nbird2id = {b: i for i, b in enumerate(target_columns)}\n","metadata":{"ExecuteTime":{"end_time":"2024-05-02T08:55:35.624574Z","start_time":"2024-05-02T08:55:35.567584Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"\ndef normalize_melspec(X, eps=1e-6):\n    mean = X.mean((1, 2), keepdim=True)\n    std = X.std((1, 2), keepdim=True)\n    Xstd = (X - mean) / (std + eps)\n\n    norm_min, norm_max = (\n        Xstd.min(-1)[0].min(-1)[0],\n        Xstd.max(-1)[0].max(-1)[0],\n    )\n    fix_ind = (norm_max - norm_min) > eps * torch.ones_like(\n        (norm_max - norm_min)\n    )\n    V = torch.zeros_like(Xstd)\n    if fix_ind.sum():\n        V_fix = Xstd[fix_ind]\n        norm_max_fix = norm_max[fix_ind, None, None]\n        norm_min_fix = norm_min[fix_ind, None, None]\n        V_fix = torch.max(\n            torch.min(V_fix, norm_max_fix),\n            norm_min_fix,\n        )\n        V_fix = (V_fix - norm_min_fix) / (norm_max_fix - norm_min_fix)\n        V[fix_ind] = V_fix\n    return V\n\n\ndef read_wav(path):\n    wav, org_sr = torchaudio.load(path, normalize=True)\n    wav = torchaudio.functional.resample(wav, orig_freq=org_sr, new_freq=mel_spec_params[\"sample_rate\"])\n    return wav\n\n\ndef crop_start_wav(wav, duration_):\n    while wav.size(-1) < duration_:\n        wav = torch.cat([wav, wav], dim=1)\n    wav = wav[:, :duration_]\n    return wav\n\n\nclass BirdDataset(torch.utils.data.Dataset):\n    def __init__(self, df, transform=None, add_secondary_labels=True):\n        self.df = df\n        self.bird2id = bird2id\n        self.num_classes = num_classes\n        self.secondary_coef = secondary_coef\n        self.add_secondary_labels = add_secondary_labels\n        self.mel_transform = torchaudio.transforms.MelSpectrogram(**mel_spec_params)\n        self.db_transform = torchaudio.transforms.AmplitudeToDB(stype='power', top_db=top_db)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def prepare_target(self, primary_label, secondary_labels):\n        secondary_labels = eval(secondary_labels)\n        target = np.zeros(self.num_classes, dtype=np.float32)\n        if primary_label != 'nocall':\n            primary_label = self.bird2id[primary_label]\n            target[primary_label] = 1.0\n            if self.add_secondary_labels:\n                for s in secondary_labels:\n                    if s != \"\" and s in self.bird2id.keys():\n                        target[self.bird2id[s]] = self.secondary_coef\n        target = torch.from_numpy(target).float()\n        return target\n\n    def prepare_spec(self, path):\n        wav = read_wav(path)\n        wav = crop_start_wav(wav, train_duration)\n        mel_spectrogram = normalize_melspec(self.db_transform(self.mel_transform(wav)))\n        mel_spectrogram = mel_spectrogram * 255\n        mel_spectrogram = mel_spectrogram.expand(3, -1, -1).permute(1, 2, 0).numpy()\n        return mel_spectrogram\n\n    def __getitem__(self, idx):\n        path = self.df[\"path\"].iloc[idx]\n        primary_label = self.df[\"primary_label\"].iloc[idx]\n        secondary_labels = self.df[\"secondary_labels\"].iloc[idx]\n        rating = self.df[\"rating\"].iloc[idx]\n\n        spec = self.prepare_spec(path)\n        target = self.prepare_target(primary_label, secondary_labels)\n\n        if self.transform is not None:\n            res = self.transform(image=spec)\n            spec = res['image'].astype(np.float32)\n        else:\n            spec = spec.astype(np.float32)\n\n        spec = spec.transpose(2, 0, 1)\n\n        return {\"spec\": spec, \"target\": target, 'rating': rating}\n","metadata":{"ExecuteTime":{"end_time":"2024-05-02T08:55:35.630720Z","start_time":"2024-05-02T08:55:35.625578Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"\nclass GeM(torch.nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super(GeM, self).__init__()\n        self.p = torch.nn.Parameter(torch.ones(1) * p)\n        self.eps = eps\n\n    def forward(self, x):\n        bs, ch, h, w = x.shape\n        x = torch.nn.functional.avg_pool2d(x.clamp(min=self.eps).pow(self.p), (x.size(-2), x.size(-1))).pow(\n            1.0 / self.p)\n        x = x.view(bs, ch)\n        return x\n\n\nclass CNN(torch.nn.Module):\n    def __init__(self, backbone, pretrained):\n        super().__init__()\n\n        out_indices = (3, 4)\n        self.backbone = timm.create_model(\n            backbone,\n            features_only=True,\n            pretrained=pretrained,\n            in_chans=3,\n            num_classes=num_classes,\n            out_indices=out_indices,\n        )\n        feature_dims = self.backbone.feature_info.channels()\n        print(f\"feature dims: {feature_dims}\")\n\n        self.global_pools = torch.nn.ModuleList([GeM() for _ in out_indices])\n        self.mid_features = np.sum(feature_dims)\n        self.neck = torch.nn.BatchNorm1d(self.mid_features)\n        self.head = torch.nn.Linear(self.mid_features, num_classes)\n\n    def forward(self, x):\n        ms = self.backbone(x)\n        h = torch.cat([global_pool(m) for m, global_pool in zip(ms, self.global_pools)], dim=1)\n        x = self.neck(h)\n        x = self.head(x)\n        return x\n","metadata":{"ExecuteTime":{"end_time":"2024-05-02T08:55:35.638517Z","start_time":"2024-05-02T08:55:35.630720Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loss","metadata":{}},{"cell_type":"code","source":"\nclass FocalLossBCE(torch.nn.Module):\n    def __init__(\n            self,\n            alpha: float = 0.25,\n            gamma: float = 2,\n            reduction: str = \"mean\",\n            bce_weight: float = 1.0,\n            focal_weight: float = 1.0,\n    ):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n        self.bce = torch.nn.BCEWithLogitsLoss(reduction=reduction)\n        self.bce_weight = bce_weight\n        self.focal_weight = focal_weight\n\n    def forward(self, logits, targets):\n        focall_loss = torchvision.ops.focal_loss.sigmoid_focal_loss(\n            inputs=logits,\n            targets=targets,\n            alpha=self.alpha,\n            gamma=self.gamma,\n            reduction=self.reduction,\n        )\n        bce_loss = self.bce(logits, targets)\n        return self.bce_weight * bce_loss + self.focal_weight * focall_loss\n\n\ncriterion = FocalLossBCE()","metadata":{"ExecuteTime":{"end_time":"2024-05-02T08:55:35.645087Z","start_time":"2024-05-02T08:55:35.638517Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Init Utils","metadata":{}},{"cell_type":"code","source":"def init_logger(log_file='train.log'):\n    from logging import INFO, FileHandler, Formatter, StreamHandler, getLogger\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\n\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count","metadata":{"ExecuteTime":{"end_time":"2024-05-02T08:55:35.652735Z","start_time":"2024-05-02T08:55:35.645087Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train and Val Functions","metadata":{}},{"cell_type":"code","source":"\ndef mixup(data, targets, alpha):\n    indices = torch.randperm(data.size(0))\n    data2 = data[indices]\n    targets2 = targets[indices]\n\n    lam = torch.FloatTensor([np.random.beta(alpha, alpha)])\n    data = data * lam + data2 * (1 - lam)\n    targets = targets * lam + targets2 * (1 - lam)\n\n    return data, targets\n\ndef train_one_epoch(model, loader, optimizer, scaler=None):\n    model.train()\n    losses = AverageMeter()\n    gt = []\n    preds = []\n    bar = tqdm(loader, total=len(loader))\n    for batch in bar:\n        optimizer.zero_grad()\n        spec = batch['spec']\n        target = batch['target']\n\n        spec, target = mixup(spec, target, 0.5)\n\n        spec = spec.to(device)\n        target = target.to(device)\n\n        if scaler is not None:\n            with torch.cuda.amp.autocast():\n                logits = model(spec)\n                loss = criterion(logits, target)\n            scaler.scale(loss).backward()\n            grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=max_grad_norm)\n            scaler.step(optimizer)\n            scaler.update()\n        else:\n            logits = model(spec)\n            loss = criterion(logits, target)\n            loss.backward()\n            grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=max_grad_norm)\n            optimizer.step()\n\n        losses.update(loss.item(), batch[\"spec\"].size(0))\n        bar.set_postfix(\n            loss=losses.avg,\n            grad=grad_norm.item(),\n            lr=optimizer.param_groups[0][\"lr\"]\n        )\n        gt.append(target.cpu().detach().numpy())\n        preds.append(logits.sigmoid().cpu().detach().numpy())\n    gt = np.concatenate(gt)\n    preds = np.concatenate(preds)\n    scores = calculate_competition_metrics_no_map(gt, preds, target_columns)\n\n    return scores, losses.avg\n\n\ndef valid_one_epoch(model, loader):\n    model.eval()\n    losses = AverageMeter()\n    bar = tqdm(loader, total=len(loader))\n    gt = []\n    preds = []\n\n    with torch.no_grad():\n        for batch in bar:\n            spec = batch['spec'].to(device)\n            target = batch['target'].to(device)\n\n            logits = model(spec)\n            loss = criterion(logits, target)\n\n            losses.update(loss.item(), batch[\"spec\"].size(0))\n\n            gt.append(target.cpu().detach().numpy())\n            preds.append(logits.sigmoid().cpu().detach().numpy())\n\n            bar.set_postfix(loss=losses.avg)\n\n    gt = np.concatenate(gt)\n    preds = np.concatenate(preds)\n    scores = calculate_competition_metrics(gt, preds, target_columns)\n    return scores, losses.avg\n\n","metadata":{"ExecuteTime":{"end_time":"2024-05-02T08:55:35.660572Z","start_time":"2024-05-02T08:55:35.652735Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Scheduler","metadata":{}},{"cell_type":"code","source":"# Fix Warmup Bug\nclass GradualWarmupSchedulerV2(GradualWarmupScheduler):\n    def __init__(self, optimizer, multiplier, total_epoch, after_scheduler=None):\n        super(GradualWarmupSchedulerV2, self).__init__(optimizer, multiplier, total_epoch, after_scheduler)\n    def get_lr(self):\n        if self.last_epoch > self.total_epoch:\n            if self.after_scheduler:\n                if not self.finished:\n                    self.after_scheduler.base_lrs = [base_lr * self.multiplier for base_lr in self.base_lrs]\n                    self.finished = True\n                return self.after_scheduler.get_lr()\n            return [base_lr * self.multiplier for base_lr in self.base_lrs]\n        if self.multiplier == 1.0:\n            return [base_lr * (float(self.last_epoch) / self.total_epoch) for base_lr in self.base_lrs]\n        else:\n            return [base_lr * ((self.multiplier - 1.) * self.last_epoch / self.total_epoch + 1.) for base_lr in self.base_lrs]","metadata":{"ExecuteTime":{"end_time":"2024-05-02T08:55:35.674443Z","start_time":"2024-05-02T08:55:35.668050Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Transformation Images","metadata":{}},{"cell_type":"code","source":"\ntransforms_train = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.Resize(image_size, image_size),\n    A.CoarseDropout(max_height=int(image_size * 0.375), max_width=int(image_size * 0.375), max_holes=1, p=0.7),\n    A.Normalize()\n])\n\ntransforms_val = A.Compose([\n    A.Resize(image_size, image_size),\n    A.Normalize()\n])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Scheduler Plot","metadata":{}},{"cell_type":"code","source":"model = CNN(backbone=backbone, pretrained=False)\nrcParams['figure.figsize'] = 20, 2\n\noptimizer = AdamW(model.parameters(), lr=lr_max, weight_decay=weight_decay)\nscheduler_cosine = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, cosine_epo)\nscheduler_warmup = GradualWarmupSchedulerV2(optimizer, multiplier=10, total_epoch=warmup_epo, after_scheduler=scheduler_cosine)\n\nlrs = []\nfor epoch in range(1, n_epochs):\n    scheduler_warmup.step()\n    lrs.append(optimizer.param_groups[0][\"lr\"])\n\nplt.plot(range(len(lrs)), lrs)","metadata":{"ExecuteTime":{"end_time":"2024-05-02T08:55:35.869499Z","start_time":"2024-05-02T08:55:35.674443Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef train_fold():\n    logger = init_logger(log_file=os.path.join(output_folder, exp_name, f\"{fold}.log\"))\n\n    logger.info(\"=\" * 90)\n    logger.info(f\"Fold {fold} Training\")\n    logger.info(\"=\" * 90)\n\n    trn_df = df[df['fold'] != fold].reset_index(drop=True)\n    val_df = df[df['fold'] == fold].reset_index(drop=True)\n    print(trn_df.shape)\n    logger.info(trn_df.shape)\n    logger.info(trn_df['primary_label'].value_counts())\n    logger.info(val_df.shape)\n    logger.info(val_df['primary_label'].value_counts())\n\n\n    trn_dataset = BirdDataset(df=trn_df.reset_index(drop=True), transform=transforms_train, add_secondary_labels=True)\n    v_ds = BirdDataset(df=val_df.reset_index(drop=True), transform=transforms_val, add_secondary_labels=True)\n\n\n    train_loader = torch.utils.data.DataLoader(trn_dataset, shuffle=True, batch_size=batch_size, drop_last=True, num_workers=num_workers, pin_memory=True)\n    val_loader = torch.utils.data.DataLoader(v_ds, shuffle=False, batch_size=batch_size, drop_last=False, num_workers=num_workers, pin_memory=True)\n\n\n    model = CNN(backbone=backbone, pretrained=True).to(device)\n    optimizer = AdamW(model.parameters(), lr=lr_max, weight_decay=weight_decay)\n    scheduler_cosine = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, cosine_epo)\n    scheduler_warmup = GradualWarmupSchedulerV2(optimizer, multiplier=10, total_epoch=warmup_epo, after_scheduler=scheduler_cosine)\n\n\n    scaler = torch.cuda.amp.GradScaler() if use_amp else None\n    patience = early_stopping\n    best_score = 0.0\n    n_patience = 0\n\n    for epoch in range(1, n_epochs + 1):\n        print(time.ctime(), 'Epoch:', epoch)\n\n        scheduler_warmup.step(epoch-1)\n\n        train_scores, train_losses_avg = train_one_epoch(model, train_loader, optimizer, scaler)\n        train_scores_str = metrics_to_string(train_scores, \"Train\")\n        train_info = f\"Epoch {epoch} - Train loss: {train_losses_avg:.4f}, {train_scores_str}\"\n        logger.info(train_info)\n\n        val_scores, val_losses_avg = valid_one_epoch(model, val_loader)\n        val_scores_str = metrics_to_string(val_scores, f\"Valid\")\n        val_info = f\"Epoch {epoch} - Valid loss: {val_losses_avg:.4f}, {val_scores_str}\"\n        logger.info(val_info)\n\n        val_score = val_scores[\"ROC\"]\n\n        is_better = val_score > best_score\n        best_score = max(val_score, best_score)\n\n        if is_better:\n            state = {\n                \"epoch\": epoch,\n                \"state_dict\": model.state_dict(),\n                \"best_loss\": best_score,\n                \"optimizer\": optimizer.state_dict(),\n            }\n            logger.info(\n                f\"Epoch {epoch} - Save Best Score: {best_score:.4f} Model\\n\")\n            torch.save(\n                state,\n                os.path.join(output_folder, exp_name, f\"{fold}.bin\")\n            )\n            n_patience = 0\n        else:\n            n_patience += 1\n            logger.info(\n                f\"Valid loss didn't improve last {n_patience} epochs.\\n\")\n\n        if n_patience >= patience:\n            logger.info(\n                \"Early stop, Training End.\\n\")\n            state = {\n                \"epoch\": epoch,\n                \"state_dict\": model.state_dict(),\n                \"best_loss\": best_score,\n                \"optimizer\": optimizer.state_dict(),\n            }\n            torch.save(\n                state,\n                os.path.join(output_folder, exp_name, f\"final_{fold}.bin\")\n            )\n            break\n\n    del model\n    torch.cuda.empty_cache()\n    gc.collect()","metadata":{"ExecuteTime":{"end_time":"2024-05-02T08:55:35.884741Z","start_time":"2024-05-02T08:55:35.869499Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_fold()","metadata":{},"execution_count":null,"outputs":[]}]}