{"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":91844,"databundleVersionId":11361821,"sourceType":"competition"},{"sourceId":11017235,"sourceType":"datasetVersion","datasetId":6859805},{"sourceId":11030424,"sourceType":"datasetVersion","datasetId":6869450}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%capture\n!pip install -U transformers  --no-index --find-links /kaggle/input/pip-hub","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:27:50.403791Z","iopub.execute_input":"2025-03-14T13:27:50.404109Z","iopub.status.idle":"2025-03-14T13:28:02.589437Z","shell.execute_reply.started":"2025-03-14T13:27:50.404076Z","shell.execute_reply":"2025-03-14T13:28:02.588524Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport os\nimport timm\nimport tqdm\nimport librosa\nimport numpy as np\nimport pandas as pd\nfrom timm.optim.adan import Adan\n# from timm.utils.model_ema import ModelEmaV2\nfrom timm.scheduler.cosine_lr import CosineLRScheduler\nfrom transformers import AutoModel\n\nimport random\nimport albumentations as A\n# import torch_audiomentations as Audio\nfrom albumentations.pytorch import ToTensorV2\n\nimport gc\nimport math\nimport glob\nimport dataclasses\nfrom transformers import set_seed\nfrom collections import defaultdict\nfrom sklearn.metrics import roc_auc_score\n\nimport matplotlib.pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:28:02.590397Z","iopub.execute_input":"2025-03-14T13:28:02.590627Z","iopub.status.idle":"2025-03-14T13:28:14.139105Z","shell.execute_reply.started":"2025-03-14T13:28:02.590605Z","shell.execute_reply":"2025-03-14T13:28:14.138275Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Parameters","metadata":{}},{"cell_type":"code","source":"train_csv_path = \"/kaggle/input/birdclef-2025/train.csv\"\nlabeldata_csv_path = \"/kaggle/input/custom-label-data/custom_label-data.csv\"\nunlabel_data_path = \"/kagglse/input/birdclef-2025/train_soundscapes\"\nsubmission_path = \"/kaggle/input/birdclef-2025/sample_submission.csv\"\nSEED = 3407\n\nVALID_DATA_RATE = .15\nTRAIN_MIN_LENGTH = 60\nEPOCHS = 30\nBS = 32\nLR = 3e-4\nWD = 3e-2\nNUM_WORKERS = 2\nMODEL_NAME = \"EfficientV2\"\nVALID_INTERVAL = 2\n\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:28:14.139863Z","iopub.execute_input":"2025-03-14T13:28:14.140268Z","iopub.status.idle":"2025-03-14T13:28:14.204359Z","shell.execute_reply.started":"2025-03-14T13:28:14.140234Z","shell.execute_reply":"2025-03-14T13:28:14.203370Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_transform = None\ntrain_audio_transform = None\n\ntest_transform = None\ntest_audio_transform = None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:28:14.205487Z","iopub.execute_input":"2025-03-14T13:28:14.205820Z","iopub.status.idle":"2025-03-14T13:28:14.225662Z","shell.execute_reply.started":"2025-03-14T13:28:14.205789Z","shell.execute_reply":"2025-03-14T13:28:14.224791Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@dataclasses.dataclass\nclass AudioParam:\n    SR: int=32_000\n    NFFT: int=2048\n    NMEL: int=128\n    FMAX: int=16_000\n    FMIN: int=20\n    HOP_LENGTH: int=NFFT // 4","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:28:14.226443Z","iopub.execute_input":"2025-03-14T13:28:14.226807Z","iopub.status.idle":"2025-03-14T13:28:14.243479Z","shell.execute_reply.started":"2025-03-14T13:28:14.226780Z","shell.execute_reply":"2025-03-14T13:28:14.242631Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"set_seed(SEED)\naudio_param = AudioParam()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:28:14.245956Z","iopub.execute_input":"2025-03-14T13:28:14.246178Z","iopub.status.idle":"2025-03-14T13:28:25.078077Z","shell.execute_reply.started":"2025-03-14T13:28:14.246147Z","shell.execute_reply":"2025-03-14T13:28:25.077381Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"sub_csv = pd.read_csv(submission_path)\nidx2cls = sub_csv.columns.drop(\"row_id\").tolist()\ncls2idx = {c: i for i, c in enumerate(idx2cls)}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:28:25.079777Z","iopub.execute_input":"2025-03-14T13:28:25.080373Z","iopub.status.idle":"2025-03-14T13:28:25.102121Z","shell.execute_reply.started":"2025-03-14T13:28:25.080348Z","shell.execute_reply":"2025-03-14T13:28:25.101449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def pcen(E, alpha=0.98, delta=2, r=0.5, s=0.025, eps=1e-6):   \n    M = scipy.signal.lfilter([s], [1, s - 1], E)\n    smooth = (eps + M)**(-alpha)\n    return (E * smooth + delta)**r - delta**r","metadata":{"trusted":true,"_kg_hide-input":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-03-14T13:28:25.102900Z","iopub.execute_input":"2025-03-14T13:28:25.103191Z","iopub.status.idle":"2025-03-14T13:28:25.108634Z","shell.execute_reply.started":"2025-03-14T13:28:25.103166Z","shell.execute_reply":"2025-03-14T13:28:25.107860Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TrainDataset(torch.utils.data.Dataset):\n    def __init__(\n        self, \n        df_group,\n        transform=None,\n        audio_transform=None,\n    ):\n        self.file_path = df_group[\"file_path\"].values\n        self.label_id = df_group[\"label_id\"].values[0]\n        self.end_time = df_group[\"end_time\"].values\n\n        self.transform = transform\n        self.audio_transform = audio_transform\n        self.length = len(self.file_path)\n\n    def __getitem__(self, idx):\n        fp, y, end_time = (\n            self.file_path[idx],\n            self.label_id,\n            self.end_time[idx],\n        )\n        x, sr = librosa.load(\n            fp,\n            sr=audio_param.SR,\n            offset=random.uniform(0, end_time-5) if end_time > 5 else 0,\n            duration=5,\n        )\n        if x.shape[0] < audio_param.SR * 5:\n            x = np.concatenate([x, np.zeros((audio_param.SR * 5 - x.shape[0]), dtype=x.dtype)])\n\n        if self.audio_transform is not None:\n            x = self.audio_transform(sample=x, sample_rate=audio_param.SR)\n\n        x = self.pipeline(x)\n\n        if self.transform is not None:\n            x = self.transform(image=x)[\"image\"]\n\n        return x, y\n\n    def pipeline(self, x):\n        mels = librosa.feature.melspectrogram(\n            y=x,\n            sr=audio_param.SR,\n            n_fft=audio_param.NFFT,\n            n_mels=audio_param.NMEL,\n            fmax=audio_param.FMAX,\n            fmin=audio_param.FMIN,\n            hop_length=audio_param.HOP_LENGTH,\n        )\n\n        # db_map = pcen(mels).astype(np.float32)\n\n        db_map = librosa.power_to_db(mels, ref=np.max)\n        db_map = (db_map + 80) / 80\n\n        return db_map[None]\n\n    def __len__(self):\n        return self.length","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-03-14T13:28:25.109427Z","iopub.execute_input":"2025-03-14T13:28:25.109705Z","iopub.status.idle":"2025-03-14T13:28:25.125686Z","shell.execute_reply.started":"2025-03-14T13:28:25.109684Z","shell.execute_reply":"2025-03-14T13:28:25.124890Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ValidDataset(TrainDataset):\n    def __init__(\n        self, \n        df_group,\n        transform=None,\n        audio_transform=None,\n    ):\n        self.file_path = df_group[\"file_path\"].values\n        self.label_id = df_group[\"label_id\"].values[0]\n        self.end_time = df_group[\"end_time\"].values\n\n        self.transform = transform\n        self.audio_transform = audio_transform\n\n        sample_counts = [max(t, 5) // 5 for t in self.end_time]\n        cumulative_counts = [0]\n        for cnt in sample_counts:\n            cumulative_counts.append(cumulative_counts[-1] + cnt)\n        self.cumulative_counts = cumulative_counts\n            \n        self.length = int(sum(sample_counts))\n\n    def __getitem__(self, idx):\n        audio_idx = next(i for i in range(len(self.cumulative_counts) - 1) \n                    if self.cumulative_counts[i] <= idx < self.cumulative_counts[i + 1])\n        inner_idx = idx - self.cumulative_counts[audio_idx]\n        \n        fp, y, end_time = (\n            self.file_path[audio_idx],\n            self.label_id,\n            self.end_time[audio_idx],\n        )\n        x, sr = librosa.load(\n            fp,\n            sr=audio_param.SR,\n            offset=inner_idx,\n            duration=5,\n        )\n\n        if x.shape[0] < audio_param.SR * 5:\n            x = np.concatenate([x, np.zeros((audio_param.SR * 5 - x.shape[0]), dtype=x.dtype)])\n\n        if self.audio_transform is not None:\n            x = self.audio_transform(sample=x, sample_rate=audio_param.SR)\n\n        x = self.pipeline(x)\n\n        if self.transform is not None:\n            x = self.transform(image=x)[\"image\"]\n\n        return x, y\n\n    def __len__(self):\n        return self.length","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-03-14T13:28:25.126333Z","iopub.execute_input":"2025-03-14T13:28:25.126604Z","iopub.status.idle":"2025-03-14T13:28:25.145689Z","shell.execute_reply.started":"2025-03-14T13:28:25.126583Z","shell.execute_reply":"2025-03-14T13:28:25.145003Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"data = pd.read_csv(labeldata_csv_path)\ndata[\"file_path\"] = data[\"file_path\"].apply(lambda x: x.replace(\"data/train_audio\", \"/kaggle/input/birdclef-2025/train_audio\"))\n\ntrain_data = []\nvalid_data = []\nfor group_id, group in data.groupby(\"label_id\"):\n    group = group.sample(frac=1).reset_index(drop=True)\n    group_length = len(group)\n    split = int(group_length * (1 - VALID_DATA_RATE))\n    if group[:split][\"end_time\"].sum() < TRAIN_MIN_LENGTH:\n        train_data.append(group)\n    else:\n        train_data.append(group[:split])\n    valid_data.append(group[split:])\n\ntrain_data = pd.concat(train_data).reset_index(drop=True)\nvalid_data = pd.concat(valid_data).reset_index(drop=True)","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-03-14T13:28:25.146337Z","iopub.execute_input":"2025-03-14T13:28:25.146576Z","iopub.status.idle":"2025-03-14T13:28:25.324849Z","shell.execute_reply.started":"2025-03-14T13:28:25.146555Z","shell.execute_reply":"2025-03-14T13:28:25.323989Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataloader = torch.utils.data.DataLoader(\n    torch.utils.data.ConcatDataset(\n        [\n            TrainDataset(\n                group,\n                transform=train_transform,\n                audio_transform=train_audio_transform,\n            ) \n            for group_id, group in train_data.groupby(\"label_id\")\n        ]\n    ),\n    batch_size=BS,\n    shuffle=True,\n    pin_memory=True,\n    num_workers=NUM_WORKERS,\n    prefetch_factor=None,\n)\n\nvalid_dataloader = torch.utils.data.DataLoader(\n    torch.utils.data.ConcatDataset(\n        [\n            ValidDataset(\n                group,\n                transform=test_transform,\n                audio_transform=test_audio_transform,\n            ) \n            for group_id, group in valid_data.groupby(\"label_id\")\n        ]\n    ),\n    batch_size=BS // 2,\n    shuffle=False,\n    pin_memory=True,\n    num_workers=NUM_WORKERS,\n    prefetch_factor=None,\n)","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-03-14T13:28:25.325868Z","iopub.execute_input":"2025-03-14T13:28:25.326189Z","iopub.status.idle":"2025-03-14T13:28:25.372721Z","shell.execute_reply.started":"2025-03-14T13:28:25.326159Z","shell.execute_reply":"2025-03-14T13:28:25.372130Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del data, train_data, valid_data\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:28:25.373423Z","iopub.execute_input":"2025-03-14T13:28:25.373710Z","iopub.status.idle":"2025-03-14T13:28:25.699576Z","shell.execute_reply.started":"2025-03-14T13:28:25.373676Z","shell.execute_reply":"2025-03-14T13:28:25.698780Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class EfficientNetV2(nn.Module):\n    def __init__(self, num_classes=1, pretrained=False, dropout=.0):\n        super().__init__()\n        self.backbone = timm.create_model(\n            \"timm/tf_efficientnetv2_m.in21k\",\n            in_chans=1,\n            pretrained=pretrained,\n            features_only=True,\n            drop_rate=dropout,\n            drop_path_rate=dropout,\n        )\n\n        self.head = nn.Sequential(\n            nn.Conv2d(512, num_classes, 1),\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(1),\n        )\n\n    def forward(self, x):\n        x = self.backbone(x)[-1]\n        x = self.head(x)\n\n        return x","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-03-14T13:28:25.700317Z","iopub.execute_input":"2025-03-14T13:28:25.700570Z","iopub.status.idle":"2025-03-14T13:28:25.716790Z","shell.execute_reply.started":"2025-03-14T13:28:25.700549Z","shell.execute_reply":"2025-03-14T13:28:25.715994Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EfficientNetB3(nn.Module):\n    def __init__(self, num_classes=1):\n        super().__init__()\n        self.backbone = timm.create_model(\n            \"efficientnet_b3.ra2_in1k\",\n            pretrained=False, \n            features_only=True, \n        )\n        self.backbone.load_state_dict(\n            torch.load(\"/kaggle/input/birdclef-2025-model-hub/efficientnet_b3.ra2_in1k/pytorch_model.bin\", weights_only=True),\n            strict=False,\n        )\n        self.backbone.conv_stem = nn.Conv2d(1, 40, 3, stride=2, padding=1, bias=False)\n        self.head = nn.Sequential(\n            nn.Conv2d(384, num_classes, 1),\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(1),\n        )\n\n    def forward(self, x):\n        x = self.backbone(x)[-1]\n        x = self.head(x)\n\n        return x","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-03-14T13:28:25.717557Z","iopub.execute_input":"2025-03-14T13:28:25.717821Z","iopub.status.idle":"2025-03-14T13:28:25.735260Z","shell.execute_reply.started":"2025-03-14T13:28:25.717788Z","shell.execute_reply":"2025-03-14T13:28:25.734498Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model = EfficientNetB3(len(idx2cls))\nmodel = EfficientNetV2(len(idx2cls), True, .1)\n\nmodel.to(device);","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:28:25.736012Z","iopub.execute_input":"2025-03-14T13:28:25.736205Z","iopub.status.idle":"2025-03-14T13:28:29.407777Z","shell.execute_reply.started":"2025-03-14T13:28:25.736188Z","shell.execute_reply":"2025-03-14T13:28:29.406840Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Preprocess","metadata":{}},{"cell_type":"code","source":"def get_param_groups(model, nowd_keys=()):\n    para_groups, para_groups_dbg = {}, {}\n    \n    for name, para in model.named_parameters():\n        if not para.requires_grad:\n            continue  # frozen weights\n        if len(para.shape) == 1 or name.endswith('.bias') or any(k in name for k in nowd_keys):\n            wd_scale, group_name = 0., 'no_decay'\n        else:\n            wd_scale, group_name = 1., 'decay'\n        \n        if group_name not in para_groups:\n            para_groups[group_name] = {'params': [], 'weight_decay_scale': wd_scale, 'lr_scale': 1.}\n            para_groups_dbg[group_name] = {'params': [], 'weight_decay_scale': wd_scale, 'lr_scale': 1.}\n        para_groups[group_name]['params'].append(para)\n        para_groups_dbg[group_name]['params'].append(name)\n\n    return list(para_groups.values())","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-03-14T13:28:29.408767Z","iopub.execute_input":"2025-03-14T13:28:29.409032Z","iopub.status.idle":"2025-03-14T13:28:29.414815Z","shell.execute_reply.started":"2025-03-14T13:28:29.409010Z","shell.execute_reply":"2025-03-14T13:28:29.414020Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def cal_score(y_hat, y):\n    matrix = torch.zeros(y_hat.shape)\n    matrix.scatter_(1, torch.from_numpy(y).reshape(-1, 1), 1)\n    matrix = matrix.numpy()\n    return roc_auc_score(\n        y_true=matrix.reshape(-1),\n        y_score=y_hat.reshape(-1),\n        average=\"macro\",\n        multi_class=\"ovo\",\n    )","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-03-14T13:28:29.415557Z","iopub.execute_input":"2025-03-14T13:28:29.415797Z","iopub.status.idle":"2025-03-14T13:28:29.465473Z","shell.execute_reply.started":"2025-03-14T13:28:29.415777Z","shell.execute_reply":"2025-03-14T13:28:29.464658Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, dataloader, criterion, optimizer):\n    model.train()\n    # dataloader = tqdm.tqdm(dataloader, desc=\"Train: \", disable=False)\n    losses = 0\n    for batch in dataloader:\n        x, y = batch\n        x, y = x.to(device), y.to(device)\n        out = model(x)\n        loss = criterion(out, y)\n        losses += loss.item()\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        \n    return losses / len(dataloader)","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-03-14T13:28:29.466191Z","iopub.execute_input":"2025-03-14T13:28:29.466389Z","iopub.status.idle":"2025-03-14T13:28:29.480852Z","shell.execute_reply.started":"2025-03-14T13:28:29.466372Z","shell.execute_reply":"2025-03-14T13:28:29.480010Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef valid_one_epoch(model, dataloader):\n    model.eval()\n    # dataloader = tqdm.tqdm(dataloader, desc=\"Valid: \", disable=False)\n    target = []\n    predict = []\n    for i, batch in enumerate(dataloader):\n        x, y = batch\n        x = x.to(device)\n        out = model(x).detach().cpu().sigmoid().numpy()\n        y = y.cpu().numpy()\n        predict.append(out)\n        target.append(y)\n\n    score = cal_score(\n        np.concatenate(predict).reshape(-1, len(idx2cls)),\n        np.concatenate(target).reshape(-1, 1),\n    )\n\n    return score","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-03-14T13:28:29.483354Z","iopub.execute_input":"2025-03-14T13:28:29.483581Z","iopub.status.idle":"2025-03-14T13:28:29.497593Z","shell.execute_reply.started":"2025-03-14T13:28:29.483553Z","shell.execute_reply":"2025-03-14T13:28:29.496704Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Criterion(nn.BCEWithLogitsLoss):\n    def __init__(self, normalize_targets=True, **kwargs):\n        super().__init__(**kwargs)\n        self.reduction = \"none\"\n        \n        self._normalize_targets = normalize_targets\n        self._eps = torch.finfo(torch.float32).eps\n        \n    def forward(self, x, y):\n        b = x.shape[0]\n        m = torch.zeros_like(x, device=x.device)\n        for i in range(b):\n            m[i, y[i]] = 1.\n\n        if self._normalize_targets:\n            m /= self._eps + m.sum(dim=1, keepdim=True)\n        per_sample_per_target_loss = -m * F.log_softmax(x, -1)\n        per_sample_loss = torch.sum(per_sample_per_target_loss, -1).mean()\n\n        return per_sample_loss","metadata":{"trusted":true,"jupyter":{"source_hidden":true},"execution":{"iopub.status.busy":"2025-03-14T13:28:29.498782Z","iopub.execute_input":"2025-03-14T13:28:29.499047Z","iopub.status.idle":"2025-03-14T13:28:29.515618Z","shell.execute_reply.started":"2025-03-14T13:28:29.499016Z","shell.execute_reply":"2025-03-14T13:28:29.514973Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = Criterion()\noptimizer = torch.optim.AdamW(\n    get_param_groups(model, (\".bias\", )),\n    lr=LR,\n    weight_decay=WD,\n)\n\nlr_scheduler = CosineLRScheduler(\n    optimizer,\n    EPOCHS,\n    warmup_t=EPOCHS // 6,\n    warmup_lr_init=LR / 10,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:28:29.516291Z","iopub.execute_input":"2025-03-14T13:28:29.516541Z","iopub.status.idle":"2025-03-14T13:28:29.533412Z","shell.execute_reply.started":"2025-03-14T13:28:29.516522Z","shell.execute_reply":"2025-03-14T13:28:29.532586Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Trainer","metadata":{}},{"cell_type":"code","source":"min_loss = float(\"inf\")\nscore = 0\n\nfor epoch in range(1, EPOCHS+1):\n    lr_scheduler.step(epoch-1)\n    loss = train_one_epoch(model, train_dataloader, criterion, optimizer)\n    if min_loss > loss:\n        min_loss = loss\n        torch.save(\n            model.state_dict(),\n            f\"{MODEL_NAME}_min-loss\"\n        )\n        \n    s = -0.01\n    if epoch % VALID_INTERVAL == 0:\n        s = valid_one_epoch(model, valid_dataloader)\n        if s > score:\n            score = s\n            torch.save(\n                model.state_dict(),\n                f\"{MODEL_NAME}_max-score\"\n            )\n\n    print(\n        f\"Epoch {epoch}/{EPOCHS}\\t\"\n        f\"Cur Epoch Loss: {loss :.4f}\\t\"\n        f\"Min Loss: {min_loss :.4f}\\t\"\n        f\"Cur Epoch Score: {s*100 :.2f}%\\t\"\n        f\"Max Score: {score*100 :.2f}%\\t\"\n    )\n\n    torch.save(\n        model.state_dict(),\n        f\"{MODEL_NAME}_last\"\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-14T13:28:29.534245Z","iopub.execute_input":"2025-03-14T13:28:29.534539Z","execution_failed":"2025-03-14T13:29:00.497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}