{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Summary of this notebook\n\nIn this notebook, I will show you how to inference with the nocall detector for train_short_audio.\n\n# Input & Output\n\n[input]\n\nbirdclef-2021 (original data)\n\n7sec clip melspectrogram images of train_short_audio \n\n(generated by kkiller's notebook https://www.kaggle.com/kneroma/birdclef-mels-computer-public)\n\nhttps://www.kaggle.com/kneroma/kkiller-birdclef-mels-computer-d7-part1\n\nhttps://www.kaggle.com/kneroma/kkiller-birdclef-mels-computer-d7-part2\n\nhttps://www.kaggle.com/kneroma/kkiller-birdclef-mels-computer-d7-part3\n\nhttps://www.kaggle.com/kneroma/kkiller-birdclef-mels-computer-d7-part4\n\nnocall detector models\n\n[output]\n\ninference results for train_short_audio are outputted.","metadata":{"papermill":{"duration":0.011818,"end_time":"2021-06-03T09:54:08.992514","exception":false,"start_time":"2021-06-03T09:54:08.980696","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import torch\n\nclass CFG:\n    debug = False\n    print_freq=100\n    num_workers=4\n    model_name= 'resnext50_32x4d'\n    dim=(128, 281)\n    epochs=10\n    batch_size=1\n    seed=42\n    target_size=2\n    fold = 3 #choose from [0,1,2,3,4]\n    pretrained = False\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"id":"alert-shopping","papermill":{"duration":1.630142,"end_time":"2021-06-03T09:54:10.633509","exception":false,"start_time":"2021-06-03T09:54:09.003367","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-20T13:42:36.235787Z","iopub.execute_input":"2023-06-20T13:42:36.236790Z","iopub.status.idle":"2023-06-20T13:42:36.244448Z","shell.execute_reply.started":"2023-06-20T13:42:36.236753Z","shell.execute_reply":"2023-06-20T13:42:36.242538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install --quiet timm\n\nimport os\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn import preprocessing\nfrom sklearn.metrics import accuracy_score, confusion_matrix\nfrom sklearn.model_selection import StratifiedKFold\n\nfrom tqdm.auto import tqdm\nfrom functools import partial\n\nimport cv2\nfrom PIL import Image\n\nimport cv2\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD\nimport torchvision.models as models\nfrom torch.nn.parameter import Parameter\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\n\nfrom albumentations.pytorch.transforms import ToTensorV2\nfrom albumentations import ImageOnlyTransform\n\nimport timm\n\nimport warnings \nwarnings.filterwarnings('ignore')\n\nimport glob","metadata":{"id":"fiscal-watson","outputId":"9d671ccb-631f-43ae-fac7-75ca3e5d0858","papermill":{"duration":11.267049,"end_time":"2021-06-03T09:54:21.911309","exception":false,"start_time":"2021-06-03T09:54:10.644260","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-20T13:42:36.249398Z","iopub.execute_input":"2023-06-20T13:42:36.250070Z","iopub.status.idle":"2023-06-20T13:42:49.599964Z","shell.execute_reply.started":"2023-06-20T13:42:36.250038Z","shell.execute_reply":"2023-06-20T13:42:49.598821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"OUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)","metadata":{"id":"genuine-belfast","papermill":{"duration":0.018627,"end_time":"2021-06-03T09:54:21.941155","exception":false,"start_time":"2021-06-03T09:54:21.922528","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-20T13:42:49.601692Z","iopub.execute_input":"2023-06-20T13:42:49.602089Z","iopub.status.idle":"2023-06-20T13:42:49.607797Z","shell.execute_reply.started":"2023-06-20T13:42:49.602051Z","shell.execute_reply":"2023-06-20T13:42:49.606936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_score(y_true, y_pred):\n    return accuracy_score(y_true, y_pred)\n\ndef get_confusion_matrix(y_true, y_pred):\n    return confusion_matrix(y_true, y_pred)\n\n@contextmanager\ndef timer(name):\n    t0 = time.time()\n    LOGGER.info(f'[{name}] start')\n    yield\n    LOGGER.info(f'[{name}] done in {time.time() - t0:.0f} s.')\n\n\ndef init_logger(log_file=OUTPUT_DIR+'train.log'):\n    from logging import getLogger, INFO, FileHandler,  Formatter,  StreamHandler\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\nLOGGER = init_logger()\n\n\ndef seed_torch(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.backends.cudnn.deterministic = True\n\nseed_torch(seed=CFG.seed)","metadata":{"id":"normal-expansion","papermill":{"duration":0.026771,"end_time":"2021-06-03T09:54:21.978393","exception":false,"start_time":"2021-06-03T09:54:21.951622","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-20T13:42:49.610588Z","iopub.execute_input":"2023-06-20T13:42:49.611688Z","iopub.status.idle":"2023-06-20T13:42:49.629143Z","shell.execute_reply.started":"2023-06-20T13:42:49.611657Z","shell.execute_reply":"2023-06-20T13:42:49.628231Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"short = pd.read_csv('../input/birdclef-2021/train_metadata.csv')","metadata":{"id":"compact-swift","papermill":{"duration":0.359205,"end_time":"2021-06-03T09:54:22.348046","exception":false,"start_time":"2021-06-03T09:54:21.988841","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-20T13:42:49.630605Z","iopub.execute_input":"2023-06-20T13:42:49.630960Z","iopub.status.idle":"2023-06-20T13:42:50.068274Z","shell.execute_reply.started":"2023-06-20T13:42:49.630930Z","shell.execute_reply":"2023-06-20T13:42:50.067285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.filenames = df['filename'].values\n        #self.labels = df['hasbird'].values\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_name = self.filenames[idx]\n        filepath = glob.glob(f'../input/birdclef2021augmentedaudio-melspec-p*/audio_images/*/{file_name}.npy')[0]\n        image = np.load(filepath)\n        image = np.stack((image,)*3, -1)\n        augmented_images = []\n        if self.transform:\n            for i in range(image.shape[0]):\n                oneimage = image[i]\n                augmented = self.transform(image=oneimage)\n                oneimage = augmented['image']\n                augmented_images.append(oneimage)\n        #label = torch.tensor(self.labels[idx]).long()\n        return np.stack(augmented_images, axis=0)#, label","metadata":{"id":"noted-chamber","papermill":{"duration":0.020988,"end_time":"2021-06-03T09:54:22.380112","exception":false,"start_time":"2021-06-03T09:54:22.359124","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-20T13:42:50.069939Z","iopub.execute_input":"2023-06-20T13:42:50.070312Z","iopub.status.idle":"2023-06-20T13:42:50.078929Z","shell.execute_reply.started":"2023-06-20T13:42:50.070276Z","shell.execute_reply":"2023-06-20T13:42:50.077926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\ndef get_transforms(*, data):\n    \n    if data == 'train':\n        return A.Compose([\n            A.Resize(CFG.dim[0], CFG.dim[1]),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.augmentations.transforms.JpegCompression(p=0.5),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])\n\n    elif data == 'valid':\n        return A.Compose([\n            A.Resize(CFG.dim[0], CFG.dim[1]),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])","metadata":{"id":"allied-negative","papermill":{"duration":0.020732,"end_time":"2021-06-03T09:54:22.411364","exception":false,"start_time":"2021-06-03T09:54:22.390632","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-20T13:42:50.080690Z","iopub.execute_input":"2023-06-20T13:42:50.081052Z","iopub.status.idle":"2023-06-20T13:42:50.093650Z","shell.execute_reply.started":"2023-06-20T13:42:50.081021Z","shell.execute_reply":"2023-06-20T13:42:50.092693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\nimport timm\n\nclass CustomResNext(nn.Module):\n    def __init__(self, model_name='resnext50_32x4d', pretrained=True):\n        super().__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        n_features = self.model.fc.in_features\n        self.model.fc = nn.Linear(n_features, CFG.target_size)\n\n    def forward(self, x):\n        x = self.model(x)\n        return x","metadata":{"id":"powered-harbor","papermill":{"duration":0.018337,"end_time":"2021-06-03T09:54:22.440037","exception":false,"start_time":"2021-06-03T09:54:22.421700","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-20T13:42:50.095021Z","iopub.execute_input":"2023-06-20T13:42:50.095433Z","iopub.status.idle":"2023-06-20T13:42:50.104918Z","shell.execute_reply.started":"2023-06-20T13:42:50.095378Z","shell.execute_reply":"2023-06-20T13:42:50.103976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def inference(model, states, test_loader, device):\n    model.to(device)\n    tk0 = tqdm(enumerate(test_loader), total=len(test_loader))\n    probs = []\n    count = 0\n    for i, (images) in tk0:\n        images = images[0]\n        images = images.to(device)\n        avg_preds = []\n        for state in states:\n            model.load_state_dict(state['model'])\n            model.eval()\n            with torch.no_grad():\n                y_preds = model(images)\n            avg_preds.append(y_preds.softmax(1).to('cpu').numpy())\n        avg_preds = np.mean(avg_preds, axis=0)\n        probs.append(avg_preds)\n        count += 1\n        if count % 100 == 0:\n            print(count)\n    #probs = np.concatenate(probs)\n    return np.asarray(probs)","metadata":{"id":"fuzzy-peoples","papermill":{"duration":0.019456,"end_time":"2021-06-03T09:54:22.470017","exception":false,"start_time":"2021-06-03T09:54:22.450561","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-20T13:42:50.106740Z","iopub.execute_input":"2023-06-20T13:42:50.107074Z","iopub.status.idle":"2023-06-20T13:42:50.119240Z","shell.execute_reply.started":"2023-06-20T13:42:50.107045Z","shell.execute_reply":"2023-06-20T13:42:50.118278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.debug == True:\n    short = short.sample(n=10)","metadata":{"id":"defensive-repeat","papermill":{"duration":0.016992,"end_time":"2021-06-03T09:54:22.497480","exception":false,"start_time":"2021-06-03T09:54:22.480488","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-20T13:42:50.122484Z","iopub.execute_input":"2023-06-20T13:42:50.122841Z","iopub.status.idle":"2023-06-20T13:42:50.129887Z","shell.execute_reply.started":"2023-06-20T13:42:50.122806Z","shell.execute_reply":"2023-06-20T13:42:50.128884Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODEL_DIR = '../input/clef-nocall-2class2-5fold/'\nmodel = CustomResNext(CFG.model_name, pretrained=CFG.pretrained)\nstates = [torch.load(MODEL_DIR+f'{CFG.model_name}_fold{CFG.fold}_best.pth'),]\ntest_dataset = TestDataset(short, transform=get_transforms(data='valid'))\ntest_loader = DataLoader(test_dataset, batch_size=CFG.batch_size, shuffle=False, \n                         num_workers=CFG.num_workers, pin_memory=True)\npredictions = inference(model, states, test_loader, CFG.device) #鳥がいる確率の方を残している(はず...)","metadata":{"id":"stopped-founder","outputId":"a13ad491-e8ac-476e-d653-38a6f321d952","papermill":{"duration":7033.89835,"end_time":"2021-06-03T11:51:36.406169","exception":false,"start_time":"2021-06-03T09:54:22.507819","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-20T13:42:50.131072Z","iopub.execute_input":"2023-06-20T13:42:50.131941Z","iopub.status.idle":"2023-06-20T14:59:24.025490Z","shell.execute_reply.started":"2023-06-20T13:42:50.131910Z","shell.execute_reply":"2023-06-20T14:59:24.024396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = [i[:,1] for i in predictions]\npredictions = [' '.join(map(str, j.tolist())) for j in predictions]\nshort['nocalldetection'] = predictions\nshort.to_csv(f'./augmented_nocalldetection_for_shortaudio_fold{CFG.fold}.csv', index=False)","metadata":{"id":"governing-andrews","outputId":"c2a70bf2-5367-4968-a6b0-21c989a18543","papermill":{"duration":2.754362,"end_time":"2021-06-03T11:51:39.356909","exception":false,"start_time":"2021-06-03T11:51:36.602547","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-20T14:59:24.027380Z","iopub.execute_input":"2023-06-20T14:59:24.028020Z","iopub.status.idle":"2023-06-20T14:59:25.339904Z","shell.execute_reply.started":"2023-06-20T14:59:24.027978Z","shell.execute_reply":"2023-06-20T14:59:25.338919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"short.head()","metadata":{"id":"oriental-burton","outputId":"8a377298-61b4-4093-af5e-868cc754cd85","papermill":{"duration":0.231852,"end_time":"2021-06-03T11:51:39.789272","exception":false,"start_time":"2021-06-03T11:51:39.557420","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-20T14:59:25.341183Z","iopub.execute_input":"2023-06-20T14:59:25.342013Z","iopub.status.idle":"2023-06-20T14:59:25.367444Z","shell.execute_reply.started":"2023-06-20T14:59:25.341978Z","shell.execute_reply":"2023-06-20T14:59:25.366427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"criminal-heart","papermill":{"duration":0.194204,"end_time":"2021-06-03T11:51:40.195036","exception":false,"start_time":"2021-06-03T11:51:40.000832","status":"completed"},"tags":[]},"execution_count":null,"outputs":[]}]}