{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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":"gpu","dataSources":[{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"}],"dockerImageVersionId":30733,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":5,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install torchlibrosa","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:10.245768Z","iopub.execute_input":"2024-06-13T10:59:10.246211Z","iopub.status.idle":"2024-06-13T10:59:22.387415Z","shell.execute_reply.started":"2024-06-13T10:59:10.246178Z","shell.execute_reply":"2024-06-13T10:59:22.386145Z"},"trusted":true},"execution_count":122,"outputs":[{"name":"stdout","text":"Requirement already satisfied: torchlibrosa in /opt/conda/lib/python3.10/site-packages (0.1.0)\nRequirement already satisfied: numpy in /opt/conda/lib/python3.10/site-packages (from torchlibrosa) (1.26.4)\nRequirement already satisfied: librosa>=0.8.0 in /opt/conda/lib/python3.10/site-packages (from torchlibrosa) (0.10.2.post1)\nRequirement already satisfied: audioread>=2.1.9 in /opt/conda/lib/python3.10/site-packages (from librosa>=0.8.0->torchlibrosa) (3.0.1)\nRequirement already satisfied: scipy>=1.2.0 in /opt/conda/lib/python3.10/site-packages (from librosa>=0.8.0->torchlibrosa) (1.11.4)\nRequirement already satisfied: scikit-learn>=0.20.0 in /opt/conda/lib/python3.10/site-packages (from librosa>=0.8.0->torchlibrosa) (1.2.2)\nRequirement already satisfied: joblib>=0.14 in /opt/conda/lib/python3.10/site-packages (from librosa>=0.8.0->torchlibrosa) (1.4.2)\nRequirement already satisfied: decorator>=4.3.0 in /opt/conda/lib/python3.10/site-packages (from librosa>=0.8.0->torchlibrosa) (5.1.1)\nRequirement already satisfied: numba>=0.51.0 in /opt/conda/lib/python3.10/site-packages (from librosa>=0.8.0->torchlibrosa) (0.58.1)\nRequirement already satisfied: soundfile>=0.12.1 in /opt/conda/lib/python3.10/site-packages (from librosa>=0.8.0->torchlibrosa) (0.12.1)\nRequirement already satisfied: pooch>=1.1 in /opt/conda/lib/python3.10/site-packages (from librosa>=0.8.0->torchlibrosa) (1.8.1)\nRequirement already satisfied: soxr>=0.3.2 in /opt/conda/lib/python3.10/site-packages (from librosa>=0.8.0->torchlibrosa) (0.3.7)\nRequirement already satisfied: typing-extensions>=4.1.1 in /opt/conda/lib/python3.10/site-packages (from librosa>=0.8.0->torchlibrosa) (4.9.0)\nRequirement already satisfied: lazy-loader>=0.1 in /opt/conda/lib/python3.10/site-packages (from librosa>=0.8.0->torchlibrosa) (0.3)\nRequirement already satisfied: msgpack>=1.0 in /opt/conda/lib/python3.10/site-packages (from librosa>=0.8.0->torchlibrosa) (1.0.7)\nRequirement already satisfied: llvmlite<0.42,>=0.41.0dev0 in /opt/conda/lib/python3.10/site-packages (from numba>=0.51.0->librosa>=0.8.0->torchlibrosa) (0.41.1)\nRequirement already satisfied: platformdirs>=2.5.0 in /opt/conda/lib/python3.10/site-packages (from pooch>=1.1->librosa>=0.8.0->torchlibrosa) (3.11.0)\nRequirement already satisfied: packaging>=20.0 in /opt/conda/lib/python3.10/site-packages (from pooch>=1.1->librosa>=0.8.0->torchlibrosa) (21.3)\nRequirement already satisfied: requests>=2.19.0 in /opt/conda/lib/python3.10/site-packages (from pooch>=1.1->librosa>=0.8.0->torchlibrosa) (2.32.3)\nRequirement already satisfied: threadpoolctl>=2.0.0 in /opt/conda/lib/python3.10/site-packages (from scikit-learn>=0.20.0->librosa>=0.8.0->torchlibrosa) (3.2.0)\nRequirement already satisfied: cffi>=1.0 in /opt/conda/lib/python3.10/site-packages (from soundfile>=0.12.1->librosa>=0.8.0->torchlibrosa) (1.16.0)\nRequirement already satisfied: pycparser in /opt/conda/lib/python3.10/site-packages (from cffi>=1.0->soundfile>=0.12.1->librosa>=0.8.0->torchlibrosa) (2.21)\nRequirement already satisfied: pyparsing!=3.0.5,>=2.0.2 in /opt/conda/lib/python3.10/site-packages (from packaging>=20.0->pooch>=1.1->librosa>=0.8.0->torchlibrosa) (3.1.1)\nRequirement already satisfied: charset-normalizer<4,>=2 in /opt/conda/lib/python3.10/site-packages (from requests>=2.19.0->pooch>=1.1->librosa>=0.8.0->torchlibrosa) (3.3.2)\nRequirement already satisfied: idna<4,>=2.5 in /opt/conda/lib/python3.10/site-packages (from requests>=2.19.0->pooch>=1.1->librosa>=0.8.0->torchlibrosa) (3.6)\nRequirement already satisfied: urllib3<3,>=1.21.1 in /opt/conda/lib/python3.10/site-packages (from requests>=2.19.0->pooch>=1.1->librosa>=0.8.0->torchlibrosa) (1.26.18)\nRequirement already satisfied: certifi>=2017.4.17 in /opt/conda/lib/python3.10/site-packages (from requests>=2.19.0->pooch>=1.1->librosa>=0.8.0->torchlibrosa) (2024.2.2)\n","output_type":"stream"}]},{"cell_type":"code","source":"from IPython.display import Audio\nimport gc\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchaudio as ta\nimport soundfile\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport os\nimport librosa\n\nimport torchvision.transforms as transforms\nfrom torch.utils.data import DataLoader,Dataset\nimport timm\nfrom pydub import AudioSegment\nimport pandas\nimport os\n\nimport torch.nn.functional as F\nfrom torchlibrosa.stft import LogmelFilterBank, Spectrogram\nfrom torchlibrosa.augmentation import SpecAugmentation\nfrom catalyst.core import Callback, CallbackOrder, IRunner\nfrom catalyst.runners.supervised import Runner, SupervisedRunner\n\nfrom sklearn import model_selection\nfrom sklearn import metrics\nimport math\nimport random\nimport json\n\nfrom transformers import AutoProcessor, AutoModelForAudioClassification","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:22.389901Z","iopub.execute_input":"2024-06-13T10:59:22.390341Z","iopub.status.idle":"2024-06-13T10:59:22.400368Z","shell.execute_reply.started":"2024-06-13T10:59:22.390303Z","shell.execute_reply":"2024-06-13T10:59:22.399299Z"},"trusted":true},"execution_count":123,"outputs":[]},{"cell_type":"code","source":"# os.environ['CUDA_LAUNCH_BLOCKING'] = '1'\n# os.environ['CUDA_VISIBLE_DEVICES']='0,1'","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:22.401591Z","iopub.execute_input":"2024-06-13T10:59:22.401896Z","iopub.status.idle":"2024-06-13T10:59:22.414365Z","shell.execute_reply.started":"2024-06-13T10:59:22.40187Z","shell.execute_reply":"2024-06-13T10:59:22.413225Z"},"trusted":true},"execution_count":124,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    ######################\n    # Globals #\n    ######################\n    epochs = 2\n    batch_size=32\n    exp_name='train7'\n    test=False\n    folds = [0]\n   \n    main_metric = \"epoch_auc_at_05\"\n\n    ######################\n    # Data #\n    ######################\n#     train_datadir = Path(\"../input/birdclef-2021/train_short_audio\") # 此处train_datadir已经被定义为path，以后在初始化中就可以直接用\"/\"表示目录\n    train_csv = \"/kaggle/input/birdclef-2024/train_metadata.csv\"\n\n    ######################\n    # Dataset #\n    ######################\n    hop_length = 715\n    n_mels = 128\n    fmin = 1000\n    sample_rate = 32000\n    fmax = 10000\n    n_fft = 1024\n\n    ######################\n    # Split #\n    ######################\n    split = \"StratifiedKFold\"\n    split_params = {\n        \"n_splits\": 5,\n        \"shuffle\": True,\n        \"random_state\": 1213\n    }\n\n    ######################\n    # Model #\n    ######################\n    pooling = \"max\"\n    pretrained = False\n    ckpt_path='./log_efficientnet/train3/fold2/checkpoints/best.pth'\n    model_name='tf_efficientnet_b2'\n    \n","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:22.416818Z","iopub.execute_input":"2024-06-13T10:59:22.417208Z","iopub.status.idle":"2024-06-13T10:59:22.427855Z","shell.execute_reply.started":"2024-06-13T10:59:22.41717Z","shell.execute_reply":"2024-06-13T10:59:22.426813Z"},"trusted":true},"execution_count":125,"outputs":[]},{"cell_type":"code","source":"class Subset(Dataset):\n    def __init__(self, root_dir, label_id=None,label=None, transform=None, seconds=5, sample_rate=32000):\n        self.root_dir=root_dir\n        self.data_dir=os.listdir(self.root_dir)\n        \n        self.label_id=label_id\n        self.label=label\n        self.label_one_hot=torch.zeros((182,))\n        self.label_one_hot[label_id]=1.0\n        \n        self.seconds=seconds\n        self.sample_rate=sample_rate\n        self.length=self.seconds*self.sample_rate\n        \n        if transform:\n            self.transform=transform\n        else:\n            pass\n#             self.transform=ta.transforms.MelSpectrogram(sample_rate=32000)\n        \n\n    def __getitem__(self, index):\n        data_path=os.path.join(self.root_dir,self.data_dir[index])\n        audio,org_sample_rate = soundfile.read(data_path)\n        \n        if self.sample_rate != org_sample_rate:\n            audio = librosa.resample(audio, orig_sr=org_sample_rate, target_sr=self.sample_rate)\n            \n        length_audio=len(audio)\n        \n        if length_audio>self.length:\n            start=random.randint(0,length_audio-self.length)\n            audio=audio[start:start+self.length]\n            \n        elif length_audio<self.length:\n            temp=np.zeros((self.length,))\n            start=random.randint(0,self.length-length_audio)\n            temp[start:start+length_audio]=audio\n            audio=temp\n            \n        else:\n            pass\n        \n        if random.randint(0,1):\n            noise=np.random.randn(self.length)\n            audio+=random.uniform(0.01,0.05)*noise\n        \n        audio=librosa.feature.melspectrogram(y=audio, sr=CFG.sample_rate, n_fft=CFG.n_fft,\n                                             n_mels=CFG.n_mels, fmin=CFG.fmin, fmax=CFG.fmax,\n                                             hop_length=CFG.hop_length)\n        \n        audio=torch.Tensor(audio)# [n_mels,time]\n\n        audio=audio.unsqueeze(0)\n#         audio=self.transform(audio)\n        \n        return {\"features\":audio, \"targets\":self.label_one_hot}\n\n    def __len__(self):\n        return len(self.data_dir)\n    \nclass MyDataset(Dataset):\n    def __init__(self,subsets,label_dict):\n        super(MyDataset,self).__init__()\n        self.subsets=subsets\n        self.label_dict=label_dict\n        \n        self.lengths=[]\n        for subset in subsets:\n            self.lengths.append(len(subset))\n    \n    def __getitem__(self,index):\n        x,y=index\n        return self.subsets[x][y]\n    \n    def __len__(self):\n        return sum(self.lengths)","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:22.429286Z","iopub.execute_input":"2024-06-13T10:59:22.429578Z","iopub.status.idle":"2024-06-13T10:59:22.447179Z","shell.execute_reply.started":"2024-06-13T10:59:22.429554Z","shell.execute_reply":"2024-06-13T10:59:22.445996Z"},"trusted":true},"execution_count":126,"outputs":[]},{"cell_type":"code","source":"class DiversitySampler(torch.utils.data.Sampler):\n    # 实现batch内的类别平均\n    def __init__(self, dataset, num_classes=182, num_classes_per_batch=8):\n        self.dataset = dataset\n        self.num_classes = num_classes\n        self.num_classes_per_batch = num_classes_per_batch\n\n    def __iter__(self):\n        batch_indices = []\n        length=self.__len__()\n        while len(batch_indices)<length:\n            for idx in random.sample(range(self.num_classes), self.num_classes_per_batch):\n                temp=[idx,random.randint(0,self.dataset.lengths[idx]-1)]\n                batch_indices.append(temp)\n        \n        return iter(batch_indices)\n\n    def __len__(self):\n        return len(self.dataset)\n    \ndef mixup_collate_fn(batch):\n    \n    audios = []\n    labels = []\n\n    for i,data in enumerate(batch):\n        if i <len(batch)-1:\n            if random.randint(0,1):\n                alpha=random.uniform(0.5,0.9)\n                audio=alpha*batch[i]['features']+(1-alpha)*batch[i+1]['features']\n                label=alpha*batch[i]['targets']+(1-alpha)*batch[i+1]['targets']\n            else:\n                audio=data['features']\n                label=data['targets']\n        audios.append(audio)\n        # labels.append(label)\n        labels.append(label)\n        \n    audios = torch.stack(audios)\n    labels = torch.stack(labels)\n    return {'features':audios, 'targets':labels}","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:22.448514Z","iopub.execute_input":"2024-06-13T10:59:22.448829Z","iopub.status.idle":"2024-06-13T10:59:22.463705Z","shell.execute_reply.started":"2024-06-13T10:59:22.448803Z","shell.execute_reply":"2024-06-13T10:59:22.462748Z"},"trusted":true},"execution_count":127,"outputs":[]},{"cell_type":"code","source":"def get_label_dict(dataset_path):\n    labels=sorted(os.listdir(dataset_path))\n    return {label:i for i, label in enumerate(labels)}\n\ndef get_dataset(path, seconds, sample_rate):\n    \n    label_dict=get_label_dict(path)\n    \n    datasets=[]\n    for label_id,label in enumerate(label_dict):\n        subset_path=os.path.join(path,label)\n\n        subset=Subset(subset_path,label_id=label_id, label=label, seconds=seconds, sample_rate=sample_rate)\n        datasets.append(subset)\n    \n    return MyDataset(datasets,label_dict)","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:22.465192Z","iopub.execute_input":"2024-06-13T10:59:22.465525Z","iopub.status.idle":"2024-06-13T10:59:22.478198Z","shell.execute_reply.started":"2024-06-13T10:59:22.465497Z","shell.execute_reply":"2024-06-13T10:59:22.47717Z"},"trusted":true},"execution_count":128,"outputs":[]},{"cell_type":"code","source":"dataset=get_dataset('/kaggle/input/birdclef-2024/train_audio', seconds=5, sample_rate=CFG.sample_rate)\n# Audio(dataset[0,1]['features'].numpy(),rate=16000)","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:22.479542Z","iopub.execute_input":"2024-06-13T10:59:22.479862Z","iopub.status.idle":"2024-06-13T10:59:22.600599Z","shell.execute_reply.started":"2024-06-13T10:59:22.479834Z","shell.execute_reply":"2024-06-13T10:59:22.599615Z"},"trusted":true},"execution_count":129,"outputs":[]},{"cell_type":"code","source":"# mean: -0.0001\n# std: 0.0303","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:22.602022Z","iopub.execute_input":"2024-06-13T10:59:22.602753Z","iopub.status.idle":"2024-06-13T10:59:22.607228Z","shell.execute_reply.started":"2024-06-13T10:59:22.602712Z","shell.execute_reply":"2024-06-13T10:59:22.606209Z"},"trusted":true},"execution_count":130,"outputs":[]},{"cell_type":"code","source":"# https://www.kaggle.com/c/rfcx-species-audio-detection/discussion/213075\nclass BCEFocalLoss(nn.Module):\n    def __init__(self, alpha=0.25, gamma=2.0):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n\n    def forward(self, preds, targets):\n        preds=preds.logits\n        bce_loss = nn.BCEWithLogitsLoss(reduction='none')(preds, targets)\n        probas = torch.sigmoid(preds)\n        loss = targets * self.alpha * \\\n            (1. - probas)**self.gamma * bce_loss + \\\n            (1. - targets) * probas**self.gamma * bce_loss\n        loss = loss.mean()\n        return loss\n    \nclass BCELoss(nn.Module):\n    def __init__(self,):\n        super().__init__()\n\n    def forward(self, preds, targets):\n        loss = nn.BCEWithLogitsLoss()(preds, targets)\n        loss = loss.mean()\n        return loss\n\nclass CELoss(nn.Module):\n    def __init__(self,):\n        super().__init__()\n\n    def forward(self, preds, targets):\n        \n        loss = nn.CrossEntropyLoss()(preds, targets)\n        loss = loss.mean()\n        return loss\n    \nclass BCEFocal2WayLoss(nn.Module):\n    def __init__(self, weights=[1, 1], class_weights=None):\n        super().__init__()\n\n        self.focal = BCEFocalLoss()\n\n        self.weights = weights\n\n    def forward(self, input, target):\n        input_ = input[\"logit\"]\n        target = target.float()\n\n        framewise_output = input[\"framewise_logit\"]\n        clipwise_output_with_max, _ = framewise_output.max(dim=1)\n\n        loss = self.focal(input_, target)\n        aux_loss = self.focal(clipwise_output_with_max, target)\n\n        return self.weights[0] * loss + self.weights[1] * aux_loss","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:22.61141Z","iopub.execute_input":"2024-06-13T10:59:22.611822Z","iopub.status.idle":"2024-06-13T10:59:22.62459Z","shell.execute_reply.started":"2024-06-13T10:59:22.611793Z","shell.execute_reply":"2024-06-13T10:59:22.623392Z"},"trusted":true},"execution_count":131,"outputs":[]},{"cell_type":"code","source":"class AUCCallback(Callback):# callback，翻译为调用，是catalyst中不同的操作的执行单位，可以使用内建的callback，也可以自己写继承类\n    def __init__(self,\n                 input_key: str = \"targets\",\n                 output_key: str = \"logits\",\n                 prefix: str = \"auc\",\n                 threshold=0.5):\n        super().__init__(CallbackOrder.Metric)\n\n        self.input_key = input_key\n        self.output_key = output_key\n        self.prefix = prefix\n        self.threshold = threshold\n\n    def on_loader_start(self, state: IRunner):\n        self.prediction: List[np.ndarray] = []\n        self.target: List[np.ndarray] = []\n\n    def on_batch_end(self, state: IRunner):# 这是类中预先定义好的钩子函数。钩子函数，一种以事件驱动的在特定情况下触发的函数，通常不需要程序员介入就可以调用\n        # 这个钩子函数顾名思义，在每个batch结束的时候调用。\n        targ = state.batch[self.input_key].detach().cpu().numpy()>self.threshold\n   \n        out = state.batch[self.output_key]\n\n        y_pred = F.softmax(out,dim=-1).detach().cpu().numpy()\n\n        self.prediction.append(y_pred)\n        self.target.append(targ)\n#         y_pred = y_pred > self.threshold\n#         y_pred = y_pred.astype(int)\n        \n#         print()\n#         print(targ.shape)\n#         print(y_pred.shape)\n#         score = metrics.f1_score(targ, y_pred, average=\"samples\")\n        score=metrics.roc_auc_score(targ, y_pred, average='micro')\n        \n        state.batch_metrics[self.prefix] = score\n\n    def on_epoch_end(self, state: IRunner):\n      \n        y_pred = np.concatenate(self.prediction, axis=0) \n        y_true = np.concatenate(self.target, axis=0)\n        \n#         score = metrics.f1_score(y_true, y_pred, average=\"samples\")\n        score=metrics.roc_auc_score(y_true, y_pred, average='micro') # micro for checking\n#         score=metrics.roc_auc_score(y_true, y_pred, average='macro',multi_class='ovr')\n\n        state.loader_metrics[self.prefix] = score\n#         if state.is_valid_loader:\n#             state.epoch_metrics['valid' + \"_epoch_\" +\n#                                 self.prefix] = score\n#         else:\n#             state.epoch_metrics[\"train_epoch_\" + self.prefix] = score\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:22.626373Z","iopub.execute_input":"2024-06-13T10:59:22.626795Z","iopub.status.idle":"2024-06-13T10:59:22.63886Z","shell.execute_reply.started":"2024-06-13T10:59:22.626764Z","shell.execute_reply":"2024-06-13T10:59:22.637839Z"},"trusted":true},"execution_count":132,"outputs":[]},{"cell_type":"code","source":"def get_scheduler(optimizer):\n    scheduler_name = CFG.scheduler_name\n\n    if scheduler_name is None:\n        return\n    else:\n        return optim.lr_scheduler.__getattribute__(scheduler_name)(\n            optimizer, **CFG.scheduler_params) # __getattribute__()这个函数用于获取类内属性，输入为str类型，获取对应类名的属性，无需自己定义\n                                                # 如果输入str为类名，则会返回该类的初始化函数\n    \ndef get_criterion():\n    if hasattr(nn, CFG.loss_name):\n        return nn.__getattribute__(CFG.loss_name)(**CFG.loss_params)\n    elif __CRITERIONS__.get(CFG.loss_name) is not None:\n        return __CRITERIONS__[CFG.loss_name](**CFG.loss_params)\n    else:\n        raise NotImplementedError\n    \ndef get_callbacks():\n    return [\n        AUCCallback(prefix=\"auc_at_05\", threshold=0.5),\n#         AUCCallback(prefix=\"auc_at_03\", threshold=0.3),\n#         AUCCallback(prefix=\"auc_at_07\", threshold=0.7),\n        #mAPCallback()\n    ]\n\ndef get_runner(device: torch.device):\n    return SupervisedRunner()\n","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:22.640447Z","iopub.execute_input":"2024-06-13T10:59:22.640828Z","iopub.status.idle":"2024-06-13T10:59:22.654694Z","shell.execute_reply.started":"2024-06-13T10:59:22.640799Z","shell.execute_reply":"2024-06-13T10:59:22.653556Z"},"trusted":true},"execution_count":133,"outputs":[]},{"cell_type":"code","source":"def init_layer(layer):\n    nn.init.xavier_uniform_(layer.weight)\n\n    if hasattr(layer, \"bias\"):\n        if layer.bias is not None:\n            layer.bias.data.fill_(0.)\n\n\ndef init_bn(bn):\n    bn.bias.data.fill_(0.)\n    bn.weight.data.fill_(1.0)\n\n\ndef init_weights(model):\n    classname = model.__class__.__name__\n    if classname.find(\"Conv2d\") != -1:\n        nn.init.xavier_uniform_(model.weight, gain=np.sqrt(2))\n        model.bias.data.fill_(0)\n    elif classname.find(\"BatchNorm\") != -1:\n        model.weight.data.normal_(1.0, 0.02)\n        model.bias.data.fill_(0)\n    elif classname.find(\"GRU\") != -1:\n        for weight in model.parameters():\n            if len(weight.size()) > 1:\n                nn.init.orghogonal_(weight.data)\n    elif classname.find(\"Linear\") != -1:\n        model.weight.data.normal_(0, 0.01)\n        model.bias.data.zero_()\n\nclass TimmSED(nn.Module):\n    def __init__(self, base_model_name: str, pretrained=False, num_classes=182, in_channels=1):\n        super().__init__()\n        \n\n        # Spec augmenter\n        self.spec_augmenter = SpecAugmentation(time_drop_width=64, time_stripes_num=2,\n                                               freq_drop_width=8, freq_stripes_num=2) # 一种mask增强技术, time_drop_width代表着时间轴上有64个维度会被block, \n                                                                                      # 同时time_stripes_num=2代表着时间轴上会发生两次block\n\n        self.bn0 = nn.BatchNorm2d(1) # n_mels对应张量的通道数\n\n        base_model = timm.create_model(\n            base_model_name, pretrained=pretrained, in_chans=in_channels) # tf_efficientnet_b0_ns\n        layers = list(base_model.children())[:-2]\n        self.encoder = nn.Sequential(*layers)\n\n        if hasattr(base_model, \"fc\"):\n            in_features = base_model.fc.in_features\n        else:\n            in_features = base_model.classifier.in_features\n        self.fc1 = nn.Sequential(\n            nn.Linear(in_features , in_features, bias=True),\n            nn.ReLU(),\n            nn.BatchNorm1d(in_features),\n        )\n        self.filter=nn.Conv1d(28,1,kernel_size=3,padding=1)\n        self.fc2 = nn.Linear(in_features, num_classes)\n\n        self.init_weight()\n\n    def init_weight(self):\n        init_layer(self.fc1[0])\n        init_layer(self.fc2)\n        init_bn(self.bn0)\n\n    def forward(self, x):\n        # (batch_size, 1, time_steps, freq_bins)\n#         x = self.spectrogram_extractor(input)\n#         x = self.logmel_extractor(x)    # (batch_size, 1, time_steps, mel_bins)\n\n        frames_num = x.shape[2] # 取时间的维度\n\n        x = self.bn0(x)\n\n        if self.training:\n            x = self.spec_augmenter(x)\n\n        # (batch_size, channels, freq, time_frames)\n        x = self.encoder(x)\n        # (batch_size, channels, frames)\n#         x = torch.mean(x, dim=2)\n        x=torch.flatten(x,start_dim=2)\n        \n        # channel smoothing\n        x1 = F.max_pool1d(x, kernel_size=3, stride=1, padding=1)\n        x2 = F.avg_pool1d(x, kernel_size=3, stride=1, padding=1)\n        x = x1 + x2\n        \n        x=x.transpose(2,1)\n        \n        x=self.filter(x)\n        \n        x=x.transpose(2,1)\n        x=x.squeeze(-1)\n#         x=x.mean(dim=-1)\n\n        x = F.dropout(x, p=0.3, training=self.training)\n     \n        x = F.relu_(self.fc1(x))\n  \n        x = F.dropout(x, p=0.3, training=self.training)\n\n        x=self.fc2(x)\n\n#         x=x.sum(dim=2)\n#         x=F.sigmoid(x)\n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:22.65654Z","iopub.execute_input":"2024-06-13T10:59:22.656874Z","iopub.status.idle":"2024-06-13T10:59:22.677214Z","shell.execute_reply.started":"2024-06-13T10:59:22.656847Z","shell.execute_reply":"2024-06-13T10:59:22.675922Z"},"trusted":true},"execution_count":134,"outputs":[]},{"cell_type":"code","source":"# validation\nsplitter = getattr(model_selection, CFG.split)(**CFG.split_params)\n\n# data\ntrain = pandas.read_csv(CFG.train_csv)\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nsampler=DiversitySampler(dataset, num_classes_per_batch=CFG.batch_size)\ntrainloader=DataLoader(dataset=dataset, sampler=sampler, batch_size=CFG.batch_size)\n# trainloader=DataLoader(dataset=dataset, shuffle=True, batch_size=8)\ncriterion = CELoss()\n\nmodel = TimmSED(CFG.model_name, pretrained=True)\n# count_parameters(model)\n# model=create_model_with_lora(model,182,8)\n# for p in model.parameters():\n#     p.requires_grad=False\n# optimizer=torch.optim.AdamW(model.parameters(), lr=0.0001)\noptimizer=torch.optim.AdamW((p for p in model.parameters() if p.requires_grad), lr=0.0003, weight_decay=0.000003)\n\ncallbacks = get_callbacks()\nrunner = get_runner(device)","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:22.678591Z","iopub.execute_input":"2024-06-13T10:59:22.678981Z","iopub.status.idle":"2024-06-13T10:59:23.190959Z","shell.execute_reply.started":"2024-06-13T10:59:22.678922Z","shell.execute_reply":"2024-06-13T10:59:23.189817Z"},"trusted":true},"execution_count":135,"outputs":[]},{"cell_type":"code","source":"# q and feed forward only\n# Num params:  94.581426 M\n# Num params:  1.347418 M\n\n# q , k and feed forward\n# Num params:  94.581426 M\n# Num params:  1.504198 M","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:23.192334Z","iopub.execute_input":"2024-06-13T10:59:23.192681Z","iopub.status.idle":"2024-06-13T10:59:23.197133Z","shell.execute_reply.started":"2024-06-13T10:59:23.192652Z","shell.execute_reply":"2024-06-13T10:59:23.195995Z"},"trusted":true},"execution_count":136,"outputs":[]},{"cell_type":"code","source":"len(optimizer.param_groups[0]['params'])","metadata":{"execution":{"iopub.status.busy":"2024-06-13T10:59:23.198402Z","iopub.execute_input":"2024-06-13T10:59:23.198755Z","iopub.status.idle":"2024-06-13T10:59:23.209241Z","shell.execute_reply.started":"2024-06-13T10:59:23.198725Z","shell.execute_reply":"2024-06-13T10:59:23.208143Z"},"trusted":true},"execution_count":137,"outputs":[{"execution_count":137,"output_type":"execute_result","data":{"text/plain":"309"},"metadata":{}}]},{"cell_type":"code","source":"# main loop\nfor i, (trn_idx, val_idx) in enumerate(splitter.split(train, y=train[\"primary_label\"])):\n    if i not in CFG.folds:\n        continue\n#     logger.info(\"=\" * 120)\n#     logger.info(f\"Fold {i} Training\")\n#     logger.info(\"=\" * 120)\n\n    trn_df = train.loc[trn_idx, :].reset_index(drop=True)\n    val_df = train.loc[val_idx, :].reset_index(drop=True)\n\n    loaders = {\n        phase: trainloader # type: ignore\n        for phase, df_ in zip([\"train\", \"valid\"], [trn_df, val_df])\n    }\n\n    runner.train(\n        model=model,\n        criterion=criterion,\n        loaders=loaders,\n        optimizer=optimizer,\n#         scheduler=scheduler,\n        num_epochs=CFG.epochs,\n        verbose=True,\n        logdir='./log_efficientnet/experiment_2/fold'+str(i+1),\n        callbacks=callbacks,\n        valid_metric=CFG.main_metric,\n        check=True# if training then false\n    )\n#         minimize_metric=CFG.minimize_metric)\n    \ndel model, optimizer\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2024-06-13T10:59:23.21062Z","iopub.execute_input":"2024-06-13T10:59:23.211064Z","iopub.status.idle":"2024-06-13T11:01:01.234526Z","shell.execute_reply.started":"2024-06-13T10:59:23.211025Z","shell.execute_reply":"2024-06-13T11:01:01.233116Z"},"trusted":true},"execution_count":138,"outputs":[{"output_type":"display_data","data":{"text/plain":"1/2 * Epoch (train):   0%|          | 0/765 [00:00<?, ?it/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"3af1841eaf474efebf5afd8453e3b60d"}},"metadata":{}},{"name":"stdout","text":"train (1/2) loss: 6.3132321039835615 | loss/mean: 6.3132321039835615 | loss/std: 0.09123900276987526 | lr: 0.0003 | momentum: 0.9\n","output_type":"stream"},{"name":"stderr","text":"/opt/conda/lib/python3.10/site-packages/accelerate/utils/dataclasses.py:301: FutureWarning: The `TPU` of `<enum 'DistributedType'>` is deprecated and will be removed in v1.0.0. Please use the `XLA` instead.\n  warnings.warn(\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"1/2 * Epoch (valid):   0%|          | 0/765 [00:00<?, ?it/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"eb986e01e2c74057874c77c37ac6c92d"}},"metadata":{}},{"name":"stderr","text":"/opt/conda/lib/python3.10/site-packages/accelerate/utils/dataclasses.py:301: FutureWarning: The `TPU` of `<enum 'DistributedType'>` is deprecated and will be removed in v1.0.0. Please use the `XLA` instead.\n  warnings.warn(\n","output_type":"stream"},{"name":"stdout","text":"valid (1/2) loss: 5.300487995147705 | loss/mean: 5.300487995147705 | loss/std: 0.028795979626561295 | lr: 0.0003 | momentum: 0.9\n* Epoch (1/2) \n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"2/2 * Epoch (train):   0%|          | 0/765 [00:00<?, ?it/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"c56fdfb811f94464ae7e8548f07d50a3"}},"metadata":{}},{"name":"stdout","text":"train (2/2) loss: 6.251863956451416 | loss/mean: 6.251863956451416 | loss/std: 0.16957894691825146 | lr: 0.0003 | momentum: 0.9\n","output_type":"stream"},{"name":"stderr","text":"/opt/conda/lib/python3.10/site-packages/accelerate/utils/dataclasses.py:301: FutureWarning: The `TPU` of `<enum 'DistributedType'>` is deprecated and will be removed in v1.0.0. Please use the `XLA` instead.\n  warnings.warn(\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"2/2 * Epoch (valid):   0%|          | 0/765 [00:00<?, ?it/s]","application/vnd.jupyter.widget-view+json":{"version_major":2,"version_minor":0,"model_id":"d565397ea5444e439fd5b5201d149f4d"}},"metadata":{}},{"name":"stderr","text":"/opt/conda/lib/python3.10/site-packages/accelerate/utils/dataclasses.py:301: FutureWarning: The `TPU` of `<enum 'DistributedType'>` is deprecated and will be removed in v1.0.0. Please use the `XLA` instead.\n  warnings.warn(\n","output_type":"stream"},{"name":"stdout","text":"valid (2/2) loss: 5.311330954233806 | loss/mean: 5.311330954233806 | loss/std: 0.06832090799103002 | lr: 0.0003 | momentum: 0.9\n* Epoch (2/2) \nTop models:\n./log_efficientnet/experiment_2/fold1/checkpoints/model.0002.pth\t2.0000\n","output_type":"stream"},{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mNameError\u001b[0m                                 Traceback (most recent call last)","Cell \u001b[0;32mIn[138], line 32\u001b[0m\n\u001b[1;32m     17\u001b[0m     runner\u001b[38;5;241m.\u001b[39mtrain(\n\u001b[1;32m     18\u001b[0m         model\u001b[38;5;241m=\u001b[39mmodel,\n\u001b[1;32m     19\u001b[0m         criterion\u001b[38;5;241m=\u001b[39mcriterion,\n\u001b[0;32m   (...)\u001b[0m\n\u001b[1;32m     28\u001b[0m         check\u001b[38;5;241m=\u001b[39m\u001b[38;5;28;01mTrue\u001b[39;00m\u001b[38;5;66;03m# if training then false\u001b[39;00m\n\u001b[1;32m     29\u001b[0m     )\n\u001b[1;32m     30\u001b[0m \u001b[38;5;66;03m#         minimize_metric=CFG.minimize_metric)\u001b[39;00m\n\u001b[0;32m---> 32\u001b[0m \u001b[38;5;28;01mdel\u001b[39;00m model, optimizer, scheduler\n\u001b[1;32m     33\u001b[0m gc\u001b[38;5;241m.\u001b[39mcollect()\n\u001b[1;32m     34\u001b[0m torch\u001b[38;5;241m.\u001b[39mcuda\u001b[38;5;241m.\u001b[39mempty_cache()\n","\u001b[0;31mNameError\u001b[0m: name 'scheduler' is not defined"],"ename":"NameError","evalue":"name 'scheduler' is not defined","output_type":"error"}]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}