{"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":"## Problem","metadata":{"id":"IByDRoRyvzMY"}},{"cell_type":"markdown","source":"**Task [[kaggle](https://www.kaggle.com/c/reface-fake-detection)]:** recognize fake videos. You need to train the binary classifier to distinguish real videos from fake ones (the provided fake data is the result of the technologies developed in Reface).\n\n****\n\n### What I should get?\n\nIn order to complete this stage, you should meet one of 2 conditions below:\n+ either make a solution with a minimum target metric value of 0.92475\n+ or be in the top 30 of all competitors.\n\n****\n\n### Evaluation\n\nThe evaluation metric for this competition is F1-Score, average='micro'. The F1 score, commonly used in information retrieval, measures accuracy using the statistics precision p and recall r. Precision is the ratio of true positives (tp) to all predicted positives (tp + fp). Recall is the ratio of true positives to all actual positives (tp + fn).\n\nThe F1 metric weights recall and precision equally, and a good retrieval algorithm will maximize both precision and recall simultaneously. Thus, moderately good performance on both will be favored over extremely good performance on one and poor performance on the other.\n\nMore information you can find at sklearn docs:\nhttps://scikit-learn.org/stable/modules/generated/sklearn.metrics.f1_score.html\n\n****\n\n### Submission\n\nFor each filename in the test set, you must predict either this file is fake video (label 1) or this file is real video (label 0). The file should contain a header and have the following format:\n\n```\nfilename,label\n004582.mp4,1\n003603.mp4,0\n```","metadata":{"id":"cVcVuY0Uvm5m"}},{"cell_type":"markdown","source":"## Install external modules and load our data","metadata":{"id":"RneR2ZDQw-_F"}},{"cell_type":"code","source":"!pip install -qq kaggle","metadata":{"id":"kdy1ta4NefeB"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir ~/.kaggle","metadata":{"id":"Jj4aovSOejSN"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp kaggle.json ~/.kaggle/","metadata":{"id":"dmJ_iN5kemII"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! chmod 600 ~/.kaggle/kaggle.json","metadata":{"id":"m37XFWeReo0c"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!kaggle datasets download -d kryvokhyzha/refacefakedetectionimages8","metadata":{"id":"JnFxPO61ezQ0","outputId":"a8d657e3-ff0a-413b-e5bc-c04b211caac4"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!unzip refacefakedetectionimages8.zip","metadata":{"id":"VNi2KAQge5M9","outputId":"ae5b382b-7e21-47f5-e490-2493eed76722"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -qq av\n!pip install -qq torchsummary\n!pip install -qq linformer\n!pip install -qq vit_pytorch","metadata":{"id":"Lnc7ynAOfFFs","outputId":"6ebb8548-c7ae-4b66-cd00-73a62ad71ccf"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -qq albumentations==1.1.0","metadata":{"id":"6YJcqvzZhPE0"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install facenet-pytorch > /dev/null 2>&1\n!apt install zip > /dev/null 2>&1","metadata":{"id":"RP7ozXs_L5Q0","execution":{"iopub.status.busy":"2021-11-04T09:55:53.27102Z","iopub.execute_input":"2021-11-04T09:55:53.271373Z","iopub.status.idle":"2021-11-04T09:56:04.998711Z","shell.execute_reply.started":"2021-11-04T09:55:53.271331Z","shell.execute_reply":"2021-11-04T09:56:04.99776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!nvidia-smi","metadata":{"id":"TD2QCNDe193q","execution":{"iopub.status.busy":"2021-11-04T09:56:05.009282Z","iopub.execute_input":"2021-11-04T09:56:05.009597Z","iopub.status.idle":"2021-11-04T09:56:05.775799Z","shell.execute_reply.started":"2021-11-04T09:56:05.009557Z","shell.execute_reply":"2021-11-04T09:56:05.775042Z"},"outputId":"ba846097-68b5-47d2-fcc3-56595fde6a51","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from google.colab import drive\ndrive.mount('/content/drive')","metadata":{"id":"saT0S6FHS2Tl","outputId":"2b255123-1202-4a94-c486-6b21a37153d6"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls /content/drive/MyDrive/dl-creator-school/","metadata":{"id":"H5yqTD-hTMiK","outputId":"d72602d9-430e-48d8-f959-e90767d132c1"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Modules importing","metadata":{"id":"dCV4IWc34im0"}},{"cell_type":"code","source":"import os\nimport glob\nimport json\nimport cv2\nimport multiprocessing as mp\n\nimport pandas as pd\nimport numpy as np\n\nimport torch\nimport torch.nn.functional as F\nimport torchvision\n\nfrom torch import nn, optim\nfrom torch.utils.data import sampler, DataLoader, Dataset\nfrom torch.optim.lr_scheduler import MultiStepLR, CosineAnnealingLR, ReduceLROnPlateau, StepLR\nfrom torch.utils import data\nfrom torchvision import transforms, models\nfrom torchvision.models import resnet101\nfrom torchsummary import summary\n\nfrom albumentations import Normalize, Compose, Resize, CenterCrop, HorizontalFlip, Rotate, VerticalFlip, RandomCrop, Downscale, RandomBrightnessContrast, GaussianBlur, HueSaturationValue\nfrom albumentations.pytorch import ToTensorV2\n\nfrom facenet_pytorch import MTCNN, InceptionResnetV1, fixed_image_standardization, training\n\nfrom linformer import Linformer\nfrom vit_pytorch.efficient import ViT\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import f1_score\n\nfrom PIL import Image\nfrom tqdm.notebook import tqdm\n\nfrom typing import List, Dict, Tuple, Union, Optional\nfrom pathlib import Path","metadata":{"id":"J3uT46T04iT2","execution":{"iopub.status.busy":"2021-11-04T09:56:05.808174Z","iopub.execute_input":"2021-11-04T09:56:05.808814Z","iopub.status.idle":"2021-11-04T09:56:09.517821Z","shell.execute_reply.started":"2021-11-04T09:56:05.808757Z","shell.execute_reply":"2021-11-04T09:56:09.516933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n%matplotlib inline\n%config InlineBackend.figure_format = 'retina'\nplt.rcParams['figure.dpi'] = 150","metadata":{"id":"aUIGxneS_3lz","execution":{"iopub.status.busy":"2021-11-04T09:56:09.519996Z","iopub.execute_input":"2021-11-04T09:56:09.520272Z","iopub.status.idle":"2021-11-04T09:56:09.541556Z","shell.execute_reply.started":"2021-11-04T09:56:09.520244Z","shell.execute_reply":"2021-11-04T09:56:09.540902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Settings","metadata":{"id":"c1bv67Wf65ck"}},{"cell_type":"code","source":"if torch.cuda.is_available():\n    device = 'cuda:0'\n#     torch.set_default_tensor_type('torch.cuda.FloatTensor')\n#     torch.multiprocessing.set_start_method('spawn')\nelse:\n    device = 'cpu'\nprint(f'Running on device: {device}')","metadata":{"id":"KALKKU9OMLuF","execution":{"iopub.status.busy":"2021-11-04T09:56:09.542464Z","iopub.execute_input":"2021-11-04T09:56:09.543198Z","iopub.status.idle":"2021-11-04T09:56:09.548656Z","shell.execute_reply.started":"2021-11-04T09:56:09.543162Z","shell.execute_reply":"2021-11-04T09:56:09.547742Z"},"outputId":"167ff084-1e64-4710-8403-564e83495eff","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !ls ../input/reface-fake-det-faces","metadata":{"execution":{"iopub.status.busy":"2021-11-04T09:56:09.550364Z","iopub.execute_input":"2021-11-04T09:56:09.550706Z","iopub.status.idle":"2021-11-04T09:56:09.564791Z","shell.execute_reply.started":"2021-11-04T09:56:09.550665Z","shell.execute_reply":"2021-11-04T09:56:09.563982Z"},"id":"soMVvDofePMQ","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PATH2PROJECT = Path('')\nPATH2DRIVE = Path('/content/drive/MyDrive/dl-creator-school/')\n\nPATH2DATA = PATH2PROJECT / 'reface-fake-detection-result'\nPATH2TRAIN = PATH2DATA / 'train'\nPATH2TEST = PATH2DATA / 'test'\nPATH2SUBMISSIONS = Path('') / 'submissions'\nPATH2CHECKOUTS = Path('') / 'checkouts'","metadata":{"id":"Q3nOA0Y765OW","execution":{"iopub.status.busy":"2021-11-04T09:58:48.737815Z","iopub.execute_input":"2021-11-04T09:58:48.738131Z","iopub.status.idle":"2021-11-04T09:58:48.743044Z","shell.execute_reply.started":"2021-11-04T09:58:48.738101Z","shell.execute_reply":"2021-11-04T09:58:48.742338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try: PATH2SUBMISSIONS.mkdir()\nexcept: pass\ntry: PATH2CHECKOUTS.mkdir()\nexcept: pass","metadata":{"id":"ZnGPv8Si_8J1","execution":{"iopub.status.busy":"2021-11-04T09:58:49.025845Z","iopub.execute_input":"2021-11-04T09:58:49.026557Z","iopub.status.idle":"2021-11-04T09:58:49.030606Z","shell.execute_reply.started":"2021-11-04T09:58:49.02651Z","shell.execute_reply":"2021-11-04T09:58:49.030017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 42\nVAL_SIZE = 0.2","metadata":{"id":"AuwPZruG8VDD","execution":{"iopub.status.busy":"2021-11-04T09:58:49.195431Z","iopub.execute_input":"2021-11-04T09:58:49.195892Z","iopub.status.idle":"2021-11-04T09:58:49.199474Z","shell.execute_reply.started":"2021-11-04T09:58:49.195845Z","shell.execute_reply":"2021-11-04T09:58:49.198619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"N_FACES = 6\n\nBATCH_SIZE = 32\nNUM_WORKERS = mp.cpu_count()\n\nWARM_UP_EPOCHS = 5\nWARM_UP_LR = 3e-3\nFINE_TUNE_EPOCHS = 20\nFINE_TUNE_LR = 5e-4\n\nH, W = 112, 112 #224, 224\nDELTA = 10\nMEAN = [0.485, 0.456, 0.406]\nSTD = [0.229, 0.224, 0.225]\n\nTHRESHOLD = 0.5\nEPSILON = 1e-7","metadata":{"id":"ERTdqVfyL2KF","execution":{"iopub.status.busy":"2021-11-04T09:58:49.40667Z","iopub.execute_input":"2021-11-04T09:58:49.406959Z","iopub.status.idle":"2021-11-04T09:58:49.413187Z","shell.execute_reply.started":"2021-11-04T09:58:49.406929Z","shell.execute_reply":"2021-11-04T09:58:49.412492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training metadata","metadata":{"id":"nyxg-uyf60_q"}},{"cell_type":"code","source":"meta_df = pd.read_csv(PATH2DATA / 'train.csv')\nmeta_df.shape","metadata":{"id":"o2uK3RTF60v_","execution":{"iopub.status.busy":"2021-11-04T09:58:49.926292Z","iopub.execute_input":"2021-11-04T09:58:49.92713Z","iopub.status.idle":"2021-11-04T09:58:49.971053Z","shell.execute_reply.started":"2021-11-04T09:58:49.927087Z","shell.execute_reply":"2021-11-04T09:58:49.970138Z"},"outputId":"8680b560-3af2-4b5f-81ee-a7c10a079a21","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta_df.label.value_counts(normalize=True)","metadata":{"id":"1YAaHThY8xCN","execution":{"iopub.status.busy":"2021-11-04T09:58:50.273929Z","iopub.execute_input":"2021-11-04T09:58:50.274233Z","iopub.status.idle":"2021-11-04T09:58:50.283872Z","shell.execute_reply.started":"2021-11-04T09:58:50.274204Z","shell.execute_reply":"2021-11-04T09:58:50.282979Z"},"outputId":"5f57dca6-007f-4d21-9e04-754c62a721a8","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta_df['path'] = meta_df['filename'].apply(lambda x: str(PATH2TRAIN / x.split('.')[0]))","metadata":{"id":"wx0VmY5WnShr","execution":{"iopub.status.busy":"2021-11-04T09:58:50.698172Z","iopub.execute_input":"2021-11-04T09:58:50.698475Z","iopub.status.idle":"2021-11-04T09:58:51.029404Z","shell.execute_reply.started":"2021-11-04T09:58:50.698442Z","shell.execute_reply":"2021-11-04T09:58:51.028095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"meta_df.sample(n=5, random_state=SEED)","metadata":{"id":"C-NMxHPQFqJ-","execution":{"iopub.status.busy":"2021-11-04T09:58:52.17673Z","iopub.execute_input":"2021-11-04T09:58:52.17702Z","iopub.status.idle":"2021-11-04T09:58:52.193774Z","shell.execute_reply.started":"2021-11-04T09:58:52.176991Z","shell.execute_reply":"2021-11-04T09:58:52.192659Z"},"outputId":"3ccf932d-1e53-4fb5-e392-e281af30d005","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Clean data","metadata":{"id":"cL7ODvphi1ds"}},{"cell_type":"markdown","source":"### Remove corrupt videos or ones in what cannot detect any faces","metadata":{"id":"cUK5lET3jIV8"}},{"cell_type":"code","source":"meta_df = meta_df[meta_df['path'].map(lambda x: os.path.exists(x))]\nmeta_df.shape","metadata":{"id":"r42a9capi1L3","execution":{"iopub.status.busy":"2021-11-04T09:58:53.542621Z","iopub.execute_input":"2021-11-04T09:58:53.542928Z","iopub.status.idle":"2021-11-04T09:59:47.395993Z","shell.execute_reply.started":"2021-11-04T09:58:53.542895Z","shell.execute_reply":"2021-11-04T09:59:47.395272Z"},"outputId":"ba8e943f-ed0f-4480-adbc-2420e189451a","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Remove videos in which do not have enough faces","metadata":{"id":"km_CsvKTjLa1"}},{"cell_type":"code","source":"# try:\n#     valid_meta_df = pd.read_csv(PATH2PROJECT / 'trainreface' / 'valid_meta_df.csv')\n# except:\nvalid_meta_df = pd.DataFrame(columns=['filename', 'label', 'path'])\nr = []\n# for row_idx, row in tqdm(train_df.iterrows()):\nfor row_idx in tqdm(meta_df.index):\n    row = meta_df.loc[row_idx]\n    img_dir = row['path']\n    face_paths = glob.glob(f'{img_dir}/*.png')\n\n    if len(face_paths) >= 4: # Satisfy the minimum requirement for the number of faces\n        r.append(row)\n\nvalid_meta_df = valid_meta_df.append(r, ignore_index=True)\n# valid_meta_df.to_csv(PATH2PROJECT / 'trainreface' / 'valid_meta_df.csv', index=False)\nvalid_meta_df.shape","metadata":{"id":"idkLNg_4jM-J","execution":{"iopub.status.busy":"2021-11-04T09:59:47.397389Z","iopub.execute_input":"2021-11-04T09:59:47.397964Z","iopub.status.idle":"2021-11-04T10:04:39.448296Z","shell.execute_reply.started":"2021-11-04T09:59:47.397928Z","shell.execute_reply":"2021-11-04T10:04:39.447654Z"},"outputId":"9b0bfb5e-a54f-466e-90cf-690be969c2bd","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_meta_df.head()","metadata":{"id":"COz2KeISMtkB","execution":{"iopub.status.busy":"2021-11-04T10:04:39.449415Z","iopub.execute_input":"2021-11-04T10:04:39.450212Z","iopub.status.idle":"2021-11-04T10:04:39.459969Z","shell.execute_reply.started":"2021-11-04T10:04:39.450182Z","shell.execute_reply":"2021-11-04T10:04:39.459372Z"},"outputId":"b2a1895b-1c22-409d-f4b0-2b1044fa0f31","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"folders = os.listdir(PATH2TEST)\nX_test = pd.DataFrame({'path': [str(PATH2TEST/folder) for folder in folders], 'filename': folders})\nlen(X_test)","metadata":{"execution":{"iopub.status.busy":"2021-11-04T10:04:39.473926Z","iopub.execute_input":"2021-11-04T10:04:39.474473Z","iopub.status.idle":"2021-11-04T10:04:39.920931Z","shell.execute_reply.started":"2021-11-04T10:04:39.474397Z","shell.execute_reply":"2021-11-04T10:04:39.920122Z"},"id":"01XfnpmIePMb","outputId":"5d8fce84-130e-46d4-bfb8-445ff501abb5","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.read_csv(PATH2PROJECT / 'sample_submission.csv')\nsubmission.shape","metadata":{"execution":{"iopub.status.busy":"2021-11-04T10:04:39.922229Z","iopub.execute_input":"2021-11-04T10:04:39.923058Z","iopub.status.idle":"2021-11-04T10:04:39.975034Z","shell.execute_reply.started":"2021-11-04T10:04:39.923013Z","shell.execute_reply":"2021-11-04T10:04:39.974381Z"},"id":"DjkkyLTXePMb","outputId":"8f6cdb82-6e99-4b24-a5d8-d7c28d01c1e2","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission['path'] = submission['filename'].apply(lambda x: str(PATH2TEST/x.split('.')[0]))","metadata":{"execution":{"iopub.status.busy":"2021-11-04T10:04:39.975946Z","iopub.execute_input":"2021-11-04T10:04:39.976419Z","iopub.status.idle":"2021-11-04T10:04:40.128702Z","shell.execute_reply.started":"2021-11-04T10:04:39.976386Z","shell.execute_reply":"2021-11-04T10:04:40.127663Z"},"id":"Hb2wCAfrePMc","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Stratified split data on test and validation","metadata":{"id":"7L0K4m7Q9cYK"}},{"cell_type":"code","source":"X_train, X_val, y_train, y_val = train_test_split(\n    valid_meta_df['path'].to_numpy(),\n    valid_meta_df['label'].to_numpy(),\n    test_size=VAL_SIZE,\n    random_state=SEED, \n    stratify=valid_meta_df['label']\n)","metadata":{"id":"zIXi_H889cHn","execution":{"iopub.status.busy":"2021-11-04T10:05:40.50592Z","iopub.execute_input":"2021-11-04T10:05:40.506235Z","iopub.status.idle":"2021-11-04T10:05:40.585745Z","shell.execute_reply.started":"2021-11-04T10:05:40.506208Z","shell.execute_reply":"2021-11-04T10:05:40.585148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.mean(y_train), np.mean(y_val)","metadata":{"id":"frm4t0w3QX6Q","execution":{"iopub.status.busy":"2021-11-04T10:05:40.946308Z","iopub.execute_input":"2021-11-04T10:05:40.946849Z","iopub.status.idle":"2021-11-04T10:05:40.955813Z","shell.execute_reply.started":"2021-11-04T10:05:40.946795Z","shell.execute_reply":"2021-11-04T10:05:40.954914Z"},"outputId":"c4a0aa36-9755-4577-bacd-421dd88d7916","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"assert not set(X_train.tolist()) & set(X_val.tolist()), 'intersection is not empty'","metadata":{"id":"h8Wz8wPK-akC","execution":{"iopub.status.busy":"2021-11-04T10:05:41.382907Z","iopub.execute_input":"2021-11-04T10:05:41.383641Z","iopub.status.idle":"2021-11-04T10:05:41.403035Z","shell.execute_reply.started":"2021-11-04T10:05:41.383598Z","shell.execute_reply":"2021-11-04T10:05:41.401993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Helper functions","metadata":{"id":"dJYubuwOwGR6"}},{"cell_type":"code","source":"def calculate_f1(preds, labels):\n    '''\n    Parameters:\n        preds: The predictions.\n        labels: The labels.\n\n    Returns:\n        f1 score\n    '''\n    return f1_score(labels, (np.array(preds) >= THRESHOLD).astype(np.uint8), average='micro')\n\n\ndef train_the_model(\n    model,\n    criterion,\n    optimizer,\n    scheduler,\n    epochs,\n    train_dataloader,\n    val_dataloader,\n    best_val_loss=1e7,\n):\n    '''\n    Parameters:\n        model: The model needs to be trained.\n        criterion: Loss function.\n        optimizer: The optimizer.\n        epochs: The number of epochs\n        train_dataloader: The dataloader used to generate training samples.\n        val_dataloader: The dataloader used to generate validation samples.\n        best_val_loss: The initial value of the best val loss (default: 1e7.)\n\n    Returns:\n        losses: All computed losses.\n        val_losses: All computed val_losses.\n        loglosses: All computed loglosses.\n        f1_scores: All computed f1_scores.\n        val_f1_scores: All computed val_f1_scores.\n        best_val_loss: New value of the best val loss.\n        best_model_state_dict: The state_dict of the best model.\n        best_optimizer_state_dict: The state_dict of the optimizer corresponds to the best model.\n    '''\n\n    losses = np.zeros(epochs)\n    val_losses = np.zeros(epochs)\n    f1_scores = np.zeros(epochs)\n    val_f1_scores = np.zeros(epochs)\n    best_model_state_dict = None\n    best_optimizer_state_dict = None\n\n    for i in tqdm(range(epochs)):\n        batch_losses = []\n        train_pbar = tqdm(train_dataloader)\n        train_pbar.desc = f'Epoch {i+1}'\n        classifier.train()\n\n        all_labels = []\n        all_preds = []\n\n        for i_batch, sample_batched in enumerate(train_pbar):\n            # Zero gradients\n            optimizer.zero_grad()\n            \n            # Make prediction.\n            y_pred = classifier(sample_batched['faces'].to(device))\n\n            all_labels.extend(sample_batched['label'].numpy().tolist())\n            all_preds.extend(y_pred.squeeze(dim=-1).detach().cpu().numpy().tolist())\n\n            # Compute loss.\n            loss = criterion(y_pred, sample_batched['label'].to(device))\n            batch_losses.append(loss.item())\n\n            # Perform a backward pass, and update the weights.\n            loss.backward()\n            optimizer.step()\n\n            # Display some information in progress-bar.\n            train_pbar.set_postfix({\n                'loss': batch_losses[-1]\n            })\n\n        # Compute scores.\n        f1_scores[i] = calculate_f1(all_preds, all_labels)\n\n        # Compute batch loss (average).\n        losses[i] = np.array(batch_losses).mean()\n\n\n        # Compute val loss\n        val_batch_losses = []\n        val_pbar = tqdm(val_dataloader)\n        val_pbar.desc = 'Validating'\n        classifier.eval()\n\n        all_labels = []\n        all_preds = []\n\n        for i_batch, sample_batched in enumerate(val_pbar):\n            # Make prediction.            \n            y_pred = classifier(sample_batched['faces'].to(device))\n\n            all_labels.extend(sample_batched['label'].numpy().tolist())\n            all_preds.extend(y_pred.squeeze(dim=-1).detach().cpu().numpy().tolist())\n\n            # Compute val loss.\n            val_loss = criterion(y_pred, sample_batched['label'].to(device))\n            val_batch_losses.append(val_loss.item())\n\n            # Display some information in progress-bar.\n            val_pbar.set_postfix({\n                'val_loss': val_batch_losses[-1]\n            })\n\n        # Compute val scores.\n        val_f1_scores[i] = calculate_f1(all_preds, all_labels)\n\n        val_losses[i] = np.array(val_batch_losses).mean()\n        print(f'loss: {losses[i]} | val loss: {val_losses[i]} | f1: {f1_scores[i]} | val f1: {val_f1_scores[i]}')\n        \n        # step of lr scheduler\n        scheduler.step(val_losses[i])\n        \n        # Update the best values\n        if val_losses[i] < best_val_loss:\n            best_val_loss = val_losses[i]\n            \n            print('Found a better checkpoint!')\n            best_model_state_dict = classifier.state_dict()\n            best_optimizer_state_dict = optimizer.state_dict()\n            state = {\n                'state_dict': best_model_state_dict,\n                'warmup_optimizer': best_optimizer_state_dict,\n                'best_val_loss': best_val_loss,\n            }\n            torch.save(state, 'best-checkout.pth')\n            \n    return losses, val_losses, f1_scores, val_f1_scores, best_val_loss, best_model_state_dict, best_optimizer_state_dict\n\n\ndef visualize_results(\n    losses,\n    val_losses,\n    f1_scores,\n    val_f1_scores\n):\n    '''\n    Parameters:\n        losses: A list of losses.\n        val_losses: A list of val losses.\n        f1_scores: A list of f1 scores.\n        val_f1_scores: A list of val f1 scores.\n    '''\n\n    fig = plt.figure(figsize=(16, 8))\n    ax = fig.add_axes([0, 0, 1, 1])\n\n    ax.plot(np.arange(1, len(losses) + 1), losses)\n    ax.plot(np.arange(1, len(val_losses) + 1), val_losses)\n    ax.set_xlabel('epoch', fontsize='xx-large')\n    ax.set_ylabel('loss', fontsize='xx-large')\n    ax.legend(\n        ['loss', 'val loss'],\n        loc='upper right',\n        fontsize='xx-large',\n        shadow=True\n    )\n    plt.show()\n\n    fig = plt.figure(figsize=(16, 8))\n    ax = fig.add_axes([0, 0, 1, 1])\n\n    ax.plot(np.arange(1, len(f1_scores) + 1), f1_scores)\n    ax.plot(np.arange(1, len(val_f1_scores) + 1), val_f1_scores)\n    ax.set_xlabel('epoch', fontsize='xx-large')\n    ax.set_ylabel('f1 score', fontsize='xx-large')\n    ax.legend(\n        ['f1', 'val f1'],\n        loc='upper left',\n        fontsize='xx-large',\n        shadow=True\n    )\n    plt.show()","metadata":{"id":"Yz5GljK-vdZ9","execution":{"iopub.status.busy":"2021-11-04T10:05:42.281705Z","iopub.execute_input":"2021-11-04T10:05:42.282129Z","iopub.status.idle":"2021-11-04T10:05:42.310998Z","shell.execute_reply.started":"2021-11-04T10:05:42.282098Z","shell.execute_reply":"2021-11-04T10:05:42.310161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset and Dataloaders","metadata":{"id":"ExtoWn0Z5Jj3"}},{"cell_type":"code","source":"class FaceDataset(Dataset):\n    def __init__(self, img_dirs, labels, n_faces=1, preprocess=None):\n        self.img_dirs = img_dirs\n        self.labels = labels\n        self.n_faces = n_faces\n        self.preprocess = preprocess\n\n    def __len__(self):\n        return len(self.img_dirs)\n    \n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n\n        img_dir = self.img_dirs[idx]\n        label = self.labels[idx]\n        face_paths = glob.glob(f'{img_dir}/*.png')\n\n        if len(face_paths) >= self.n_faces:\n            sample = sorted(np.random.choice(face_paths, self.n_faces, replace=False))\n        else:\n            sample = sorted(np.random.choice(face_paths, self.n_faces, replace=True))\n            \n        faces = []\n        for face_path in sample:\n            face = cv2.imread(face_path, 1)\n            face = cv2.cvtColor(face, cv2.COLOR_BGR2RGB)\n            faces.append(face)\n            \n        if self.preprocess is not None:\n            d = {f'image{i-1}': faces[i] for i in range(1, self.n_faces)}\n            d['image'] = faces[0]\n            faces = list(self.preprocess(**d).values())\n\n        return {'faces': torch.stack(faces).permute(1, 0, 2, 3), 'label': torch.tensor([label], dtype=torch.float)}#{'faces': np.concatenate(faces, axis=-1).transpose(2, 0, 1), 'label': np.array([label], dtype=float)}","metadata":{"id":"Rqxg5ksv5LRG","execution":{"iopub.status.busy":"2021-11-04T10:05:43.035415Z","iopub.execute_input":"2021-11-04T10:05:43.035823Z","iopub.status.idle":"2021-11-04T10:05:43.047632Z","shell.execute_reply.started":"2021-11-04T10:05:43.035792Z","shell.execute_reply":"2021-11-04T10:05:43.046809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transforms = Compose([\n    Resize(H+DELTA, W+DELTA),\n    # Downscale(scale_min=0.5, scale_max=0.9, p=0.3),\n    RandomCrop(H, W),\n    HorizontalFlip(p=0.5),\n    RandomBrightnessContrast(brightness_limit=0, contrast_limit=0.2, p=0.3),\n    HueSaturationValue(hue_shift_limit=20, sat_shift_limit=30, val_shift_limit=20, p=0.3),\n    GaussianBlur(blur_limit=(3, 7), p=0.3),\n    Normalize(mean=MEAN, std=STD, p=1),\n    ToTensorV2()\n], additional_targets={f'image{i}': 'image' for i in range(0, N_FACES-1)})\n\nval_transforms = Compose([\n    Resize(H+DELTA, W+DELTA),\n    CenterCrop(H, W),\n    Normalize(mean=MEAN, std=STD, p=1),\n    ToTensorV2()\n], additional_targets={f'image{i}': 'image' for i in range(0, N_FACES-1)})\n\ntest_transforms = Compose([\n    Resize(H+DELTA, W+DELTA),\n    CenterCrop(H, W),\n    Normalize(mean=MEAN, std=STD, p=1),\n    ToTensorV2()\n], additional_targets={f'image{i}': 'image' for i in range(0, N_FACES-1)})","metadata":{"id":"X_MmezDumM2u","execution":{"iopub.status.busy":"2021-11-04T10:05:45.176391Z","iopub.execute_input":"2021-11-04T10:05:45.17669Z","iopub.status.idle":"2021-11-04T10:05:45.187187Z","shell.execute_reply.started":"2021-11-04T10:05:45.176653Z","shell.execute_reply":"2021-11-04T10:05:45.186334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = FaceDataset(\n    img_dirs=X_train,\n    labels=y_train,\n    n_faces=N_FACES,\n    preprocess=train_transforms\n)\nval_dataset = FaceDataset(\n    img_dirs=X_val,\n    labels=y_val,\n    n_faces=N_FACES,\n    preprocess=val_transforms\n)\ntest_dataset = FaceDataset(\n    img_dirs=X_test['path'].values,\n    labels=[0]*len(X_test['path']),\n    n_faces=N_FACES,\n    preprocess=test_transforms\n)\n\ntrain_dataloader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=NUM_WORKERS,\n#     generator=torch.Generator(device='cuda'),\n#     num_workers=0,\n#     pin_memory=False,\n)\nval_dataloader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n#     generator=torch.Generator(device='cuda'),\n#     num_workers=0,\n#     pin_memory=False,\n)\ntest_dataloader = DataLoader(\n    test_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n#     generator=torch.Generator(device='cuda'),\n#     num_workers=0,\n#     pin_memory=False,\n)","metadata":{"id":"n6vPVKHTmQOn","execution":{"iopub.status.busy":"2021-11-04T10:05:45.762292Z","iopub.execute_input":"2021-11-04T10:05:45.762573Z","iopub.status.idle":"2021-11-04T10:05:45.774306Z","shell.execute_reply.started":"2021-11-04T10:05:45.762542Z","shell.execute_reply":"2021-11-04T10:05:45.773527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10,10))\nfor ii,img in enumerate(next(iter(train_dataloader))['faces'][0].permute(1, 0, 2, 3)):\n    plt.subplot(2,2,ii+1)\n    mean = np.array([0.485, 0.456, 0.406])\n    std = np.array([0.229, 0.224, 0.225])\n    inp = img.numpy().transpose((1, 2, 0))\n    inp = std * inp + mean\n    inp = np.clip(inp, 0, 1)\n    plt.imshow(inp)\n    if ii == 3:\n        break","metadata":{"execution":{"iopub.status.busy":"2021-11-04T10:07:08.176562Z","iopub.execute_input":"2021-11-04T10:07:08.176948Z","iopub.status.idle":"2021-11-04T10:07:18.684112Z","shell.execute_reply.started":"2021-11-04T10:07:08.176912Z","shell.execute_reply":"2021-11-04T10:07:18.683429Z"},"id":"tHiROSYbePMl","outputId":"9ac1c8bf-9cf7-4b66-e72f-908a5be7fe45","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"next(iter(train_dataloader))['faces'].shape","metadata":{"id":"kjQtlPbSDbhD","execution":{"iopub.status.busy":"2021-11-04T10:05:57.270716Z","iopub.status.idle":"2021-11-04T10:05:57.271339Z","shell.execute_reply.started":"2021-11-04T10:05:57.271128Z","shell.execute_reply":"2021-11-04T10:05:57.271149Z"},"outputId":"8ff22d73-24e4-4dbd-9e7f-97f98de6b193","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Models","metadata":{"id":"tHhCvBsf5SB6"}},{"cell_type":"code","source":"class DeepfakeClassifierResnet(nn.Module):\n    def __init__(self, encoder, in_channels=3, out_channels=64, kernel_size=7, stride=2, padding=3, bias=False, linear_size=2048, num_classes=1):\n        super(DeepfakeClassifierResnet, self).__init__()\n        self.encoder = encoder\n        \n        # Modify input layer.\n        self.encoder.conv1 = nn.Conv2d(\n            in_channels,\n            out_channels,\n            kernel_size=kernel_size,\n            stride=stride,\n            padding=padding,\n            bias=bias,\n        )\n        \n        self.encoder.fc = nn.Linear(linear_size * 1, num_classes)\n\n    def forward(self, x):\n        return torch.sigmoid(self.encoder(x))\n    \n    def freeze_all_layers(self):\n        for param in self.encoder.parameters():\n            param.requires_grad = False\n\n    def freeze_middle_layers(self):\n        self.freeze_all_layers()\n        \n        for param in self.encoder.conv1.parameters():\n            param.requires_grad = True\n            \n        for param in self.encoder.fc.parameters():\n            param.requires_grad = True\n\n    def unfreeze_all_layers(self):\n        for param in self.encoder.parameters():\n            param.requires_grad = True","metadata":{"id":"KLWlt9sI5L49","execution":{"iopub.status.busy":"2021-11-04T10:07:29.229805Z","iopub.execute_input":"2021-11-04T10:07:29.230266Z","iopub.status.idle":"2021-11-04T10:07:29.241045Z","shell.execute_reply.started":"2021-11-04T10:07:29.230231Z","shell.execute_reply":"2021-11-04T10:07:29.240137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DeepfakeClassifierInception(nn.Module):\n    def __init__(self, encoder, in_channels=3, out_channels=32, kernel_size=3, stride=2, padding=3, bias=False, linear_size=512, num_classes=1):\n        super(DeepfakeClassifierInception, self).__init__()\n        self.encoder = encoder\n        \n        # Modify input layer.\n        self.encoder.conv2d_1a.conv = nn.Conv2d(\n            in_channels,\n            out_channels,\n            kernel_size=kernel_size,\n            stride=stride,\n            padding=padding,\n            bias=bias,\n        )\n        \n        self.encoder.logits = nn.Linear(linear_size * 1, num_classes)\n\n    def forward(self, x):\n        return torch.sigmoid(self.encoder(x))\n    \n    def freeze_all_layers(self):\n        for param in self.encoder.parameters():\n            param.requires_grad = False\n\n    def freeze_middle_layers(self):\n        self.freeze_all_layers()\n        \n        for param in self.encoder.conv2d_1a.conv.parameters():\n            param.requires_grad = True\n            \n        for param in self.encoder.logits.parameters():\n            param.requires_grad = True\n\n    def unfreeze_all_layers(self):\n        for param in self.encoder.parameters():\n            param.requires_grad = True","metadata":{"execution":{"iopub.status.busy":"2021-11-04T10:07:29.431299Z","iopub.execute_input":"2021-11-04T10:07:29.431737Z","iopub.status.idle":"2021-11-04T10:07:29.441818Z","shell.execute_reply.started":"2021-11-04T10:07:29.431706Z","shell.execute_reply":"2021-11-04T10:07:29.440956Z"},"id":"FaR5ErtsePMq","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DeepfakeClassifierR3D18(nn.Module):\n    def __init__(self, encoder, linear_size=512, num_classes=1):\n        super(DeepfakeClassifierR3D18, self).__init__()\n        self.encoder = encoder\n        \n        # Modify output layer.\n        num_features = self.encoder.fc.in_features\n        self.encoder.fc = nn.Linear(num_features, num_classes)\n\n    def forward(self, x):\n        return torch.sigmoid(self.encoder(x))\n    \n    def freeze_all_layers(self):\n        for param in self.encoder.parameters():\n            param.requires_grad = False\n\n    def freeze_middle_layers(self):\n        self.freeze_all_layers()\n            \n        for param in self.encoder.fc.parameters():\n            param.requires_grad = True\n\n    def unfreeze_all_layers(self):\n        for param in self.encoder.parameters():\n            param.requires_grad = True","metadata":{"execution":{"iopub.status.busy":"2021-11-04T10:07:29.640913Z","iopub.execute_input":"2021-11-04T10:07:29.64132Z","iopub.status.idle":"2021-11-04T10:07:29.6503Z","shell.execute_reply.started":"2021-11-04T10:07:29.641283Z","shell.execute_reply":"2021-11-04T10:07:29.649488Z"},"id":"IbsTK_MJePMr","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class FocalLoss(nn.Module):\n    def __init__(self, gamma=2, sample_weight=None):\n        super().__init__()\n        self.gamma = gamma\n        self.sample_weight = sample_weight\n\n    def forward(self, logit, target):\n        target = target.float()\n        max_val = (-logit).clamp(min=0)\n        loss = logit - logit * target + max_val + \\\n               ((-max_val).exp() + (-logit - max_val).exp()).log()\n\n        invprobs = F.logsigmoid(-logit * (target * 2.0 - 1.0))\n        loss = (invprobs * self.gamma).exp() * loss\n        if len(loss.size())==2:\n            loss = loss.sum(dim=1)\n        if self.sample_weight is not None:\n            loss = loss * self.sample_weight\n        return loss.mean()","metadata":{"execution":{"iopub.status.busy":"2021-11-04T10:07:29.845678Z","iopub.execute_input":"2021-11-04T10:07:29.846529Z","iopub.status.idle":"2021-11-04T10:07:29.854736Z","shell.execute_reply.started":"2021-11-04T10:07:29.84649Z","shell.execute_reply":"2021-11-04T10:07:29.854093Z"},"id":"f6M_FFIKePMt","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# efficient_transformer = Linformer(\n#     dim=128,\n#     seq_len=49+1,  # 7x7 patches + 1 cls-token\n#     depth=12,\n#     heads=8,\n#     k=64\n# )\n\n# classifier = ViT(\n#     dim=128,\n#     image_size=224,\n#     patch_size=32,\n#     num_classes=1,\n#     transformer=efficient_transformer,\n#     channels=3,\n# ).to(device)\n# classifier.train()","metadata":{"execution":{"iopub.status.busy":"2021-11-04T10:07:30.043139Z","iopub.execute_input":"2021-11-04T10:07:30.043432Z","iopub.status.idle":"2021-11-04T10:07:30.047996Z","shell.execute_reply.started":"2021-11-04T10:07:30.043402Z","shell.execute_reply":"2021-11-04T10:07:30.04698Z"},"id":"0g6bg2bzePMt","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder_r3d_18 = models.video.r3d_18(\n    pretrained=True,\n)\n\nclassifier = DeepfakeClassifierR3D18(encoder=encoder_r3d_18, linear_size=512, num_classes=1)\n\nclassifier.to(device);\nclassifier.train();","metadata":{"execution":{"iopub.status.busy":"2021-11-04T10:07:30.192821Z","iopub.execute_input":"2021-11-04T10:07:30.193128Z","iopub.status.idle":"2021-11-04T10:07:36.609556Z","shell.execute_reply.started":"2021-11-04T10:07:30.193094Z","shell.execute_reply":"2021-11-04T10:07:36.60877Z"},"id":"ubA5Bm0XePMu","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# x = torch.zeros(1, 3, N_FACES, H, W)\n# y= classifier(x)\n# print(y.shape)","metadata":{"execution":{"iopub.status.busy":"2021-11-04T10:07:36.611001Z","iopub.execute_input":"2021-11-04T10:07:36.612789Z","iopub.status.idle":"2021-11-04T10:07:36.617042Z","shell.execute_reply.started":"2021-11-04T10:07:36.612753Z","shell.execute_reply":"2021-11-04T10:07:36.615803Z"},"id":"L8lSfP3TePMv","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# encoder_facenet = InceptionResnetV1(\n#     classify=True,\n#     pretrained='casia-webface',\n#     num_classes=1\n# )\n\n# classifier = DeepfakeClassifierInception(encoder=encoder_facenet, in_channels=3*N_FACES, num_classes=1)\n\n# classifier.to(device);\n# classifier.train();","metadata":{"execution":{"iopub.status.busy":"2021-11-04T10:07:36.618472Z","iopub.execute_input":"2021-11-04T10:07:36.618832Z","iopub.status.idle":"2021-11-04T10:07:36.6281Z","shell.execute_reply.started":"2021-11-04T10:07:36.618793Z","shell.execute_reply":"2021-11-04T10:07:36.627381Z"},"id":"fF7JylFWePMv","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# encoder_resnet = resnet101(pretrained=True)\n\n# classifier = DeepfakeClassifierResnet(encoder=encoder_resnet, in_channels=3*N_FACES, num_classes=1)\n\n# classifier.to(device);\n# classifier.train();","metadata":{"id":"SJr-XcCHnua7","execution":{"iopub.status.busy":"2021-11-04T10:07:36.630288Z","iopub.execute_input":"2021-11-04T10:07:36.630626Z","iopub.status.idle":"2021-11-04T10:07:36.640499Z","shell.execute_reply.started":"2021-11-04T10:07:36.630587Z","shell.execute_reply":"2021-11-04T10:07:36.639864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.BCELoss()#FocalLoss()","metadata":{"id":"BKpiduWCnySc","execution":{"iopub.status.busy":"2021-11-04T10:07:36.641368Z","iopub.execute_input":"2021-11-04T10:07:36.641913Z","iopub.status.idle":"2021-11-04T10:07:36.652008Z","shell.execute_reply.started":"2021-11-04T10:07:36.641881Z","shell.execute_reply":"2021-11-04T10:07:36.651191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"losses = np.zeros(WARM_UP_EPOCHS + FINE_TUNE_EPOCHS)\nval_losses = np.zeros(WARM_UP_EPOCHS + FINE_TUNE_EPOCHS)\nf1_scores = np.zeros(WARM_UP_EPOCHS + FINE_TUNE_EPOCHS)\nval_f1_scores = np.zeros(WARM_UP_EPOCHS + FINE_TUNE_EPOCHS)\n\nbest_val_loss = 1e7","metadata":{"id":"Z_K0Q_O7n2Dt","execution":{"iopub.status.busy":"2021-11-04T10:07:36.653303Z","iopub.execute_input":"2021-11-04T10:07:36.653548Z","iopub.status.idle":"2021-11-04T10:07:36.665171Z","shell.execute_reply.started":"2021-11-04T10:07:36.653519Z","shell.execute_reply":"2021-11-04T10:07:36.66444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define training hyperparameters","metadata":{"id":"SFp3sZ-A5a-q"}},{"cell_type":"code","source":"classifier.freeze_middle_layers()\nwarmup_optimizer = optim.Adam(filter(lambda p: p.requires_grad, classifier.parameters()), lr=WARM_UP_LR)\nwarmup_scheduler = ReduceLROnPlateau(warmup_optimizer, mode='min', factor=0.1, patience=5, threshold=0.0001, threshold_mode='abs', verbose=True)","metadata":{"id":"95uDoQm7n-t8","execution":{"iopub.status.busy":"2021-11-04T10:07:36.666143Z","iopub.execute_input":"2021-11-04T10:07:36.666588Z","iopub.status.idle":"2021-11-04T10:07:36.68056Z","shell.execute_reply.started":"2021-11-04T10:07:36.666537Z","shell.execute_reply":"2021-11-04T10:07:36.679692Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"summary(classifier, input_size=(3, N_FACES, H, W))","metadata":{"execution":{"iopub.status.busy":"2021-11-04T10:07:36.682211Z","iopub.execute_input":"2021-11-04T10:07:36.683134Z","iopub.status.idle":"2021-11-04T10:07:38.096714Z","shell.execute_reply.started":"2021-11-04T10:07:36.683091Z","shell.execute_reply":"2021-11-04T10:07:38.095826Z"},"id":"9EZ-cqgbePMz","outputId":"13e093b3-da98-4f18-8b85-c2d543f34a80","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{"id":"GeUU3tdk5kaP"}},{"cell_type":"code","source":"losses[:WARM_UP_EPOCHS], val_losses[:WARM_UP_EPOCHS], \\\nf1_scores[:WARM_UP_EPOCHS], val_f1_scores[:WARM_UP_EPOCHS], \\\nbest_val_loss, \\\nbest_model_state_dict, best_optimizer_state_dict \\\n= train_the_model(\n    model=classifier,\n    criterion=criterion,\n    optimizer=warmup_optimizer,\n    scheduler=warmup_scheduler,\n    epochs=WARM_UP_EPOCHS,\n    train_dataloader=train_dataloader,\n    val_dataloader=val_dataloader,\n    best_val_loss=best_val_loss,\n)\n\n# Save the best checkpoint.\nif best_model_state_dict is not None:\n    state = {\n        'state_dict': best_model_state_dict,\n        'warmup_optimizer': best_optimizer_state_dict,\n        'best_val_loss': best_val_loss,\n    }\n    torch.save(state, 'best-checkout-warmup.pth')","metadata":{"id":"5Ki89wSL5kJ-","execution":{"iopub.status.busy":"2021-11-04T10:07:38.103901Z","iopub.execute_input":"2021-11-04T10:07:38.104448Z","iopub.status.idle":"2021-11-04T10:08:17.5036Z","shell.execute_reply.started":"2021-11-04T10:07:38.104413Z","shell.execute_reply":"2021-11-04T10:08:17.502334Z"},"outputId":"4a3fd097-0e95-4c6e-b515-688ebd1c1333","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"visualize_results(\n    losses=losses[:WARM_UP_EPOCHS],\n    val_losses=val_losses[:WARM_UP_EPOCHS],\n    f1_scores=f1_scores[:WARM_UP_EPOCHS],\n    val_f1_scores=val_f1_scores[:WARM_UP_EPOCHS]\n)","metadata":{"id":"RJr8GgXcoyN4","execution":{"iopub.status.busy":"2021-11-02T22:56:14.580671Z","iopub.status.idle":"2021-11-02T22:56:14.581129Z","shell.execute_reply.started":"2021-11-02T22:56:14.580881Z","shell.execute_reply":"2021-11-02T22:56:14.580909Z"},"outputId":"89d3c966-ebb9-4ac6-ba8e-e00a654bd24f","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# state = torch.load(PATH2PROJECT / 'trainreface' / 'temp-best-checkout-resnet101.pth', map_location=lambda storage, loc: storage)\nstate = torch.load('best-checkout-warmup.pth', map_location=lambda storage, loc: storage)\n# state = torch.load('best-checkout.pth', map_location=lambda storage, loc: storage)\nbest_val_loss = state['best_val_loss']\nclassifier.load_state_dict(state['state_dict'])\n# classifier.to(device)","metadata":{"execution":{"iopub.status.busy":"2021-11-02T23:01:28.445913Z","iopub.execute_input":"2021-11-02T23:01:28.446193Z","iopub.status.idle":"2021-11-02T23:01:28.574534Z","shell.execute_reply.started":"2021-11-02T23:01:28.446164Z","shell.execute_reply":"2021-11-02T23:01:28.573555Z"},"id":"b8pblvSkePM3","outputId":"f702d8a1-0651-4da7-db3a-9699cacb1368","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"classifier.unfreeze_all_layers()\nfinetune_optimizer = optim.Adam(filter(lambda p: p.requires_grad, classifier.parameters()), lr=FINE_TUNE_LR)\nfinetune_scheduler = ReduceLROnPlateau(finetune_optimizer, mode='min', factor=0.1, patience=5, threshold=0.0001, threshold_mode='abs', verbose=True)","metadata":{"id":"kxqhx7cS0t5B","execution":{"iopub.status.busy":"2021-11-02T23:01:36.275878Z","iopub.execute_input":"2021-11-02T23:01:36.276623Z","iopub.status.idle":"2021-11-02T23:01:36.284984Z","shell.execute_reply.started":"2021-11-02T23:01:36.276575Z","shell.execute_reply":"2021-11-02T23:01:36.283749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# losses[WARM_UP_EPOCHS:WARM_UP_EPOCHS+FINE_TUNE_EPOCHS], val_losses[WARM_UP_EPOCHS:WARM_UP_EPOCHS+FINE_TUNE_EPOCHS], \\\n# f1_scores[WARM_UP_EPOCHS:WARM_UP_EPOCHS+FINE_TUNE_EPOCHS], val_f1_scores[WARM_UP_EPOCHS:WARM_UP_EPOCHS+FINE_TUNE_EPOCHS], \\\n# best_val_loss, best_val_logloss, \\\n# best_model_state_dict, best_optimizer_state_dict \\\n_, _, \\\n_, _, \\\nbest_val_loss, \\\nbest_model_state_dict, best_optimizer_state_dict \\\n= train_the_model(\n    model=classifier,\n    criterion=criterion,\n    optimizer=finetune_optimizer,\n    scheduler=finetune_scheduler,\n    epochs=FINE_TUNE_EPOCHS,\n    train_dataloader=train_dataloader,\n    val_dataloader=val_dataloader,\n    best_val_loss=best_val_loss,\n)\n\n# Save the best checkpoint.\nif best_model_state_dict is not None:\n    state = {\n        'state_dict': best_model_state_dict,\n        'finetune_optimizer': best_optimizer_state_dict,\n        'best_val_loss': best_val_loss,\n    }\n\n    torch.save(state, 'best-checkout-finetune.pth')","metadata":{"execution":{"iopub.status.busy":"2021-11-03T01:47:57.270371Z","iopub.execute_input":"2021-11-03T01:47:57.270884Z","iopub.status.idle":"2021-11-03T01:49:15.871701Z","shell.execute_reply.started":"2021-11-03T01:47:57.270845Z","shell.execute_reply":"2021-11-03T01:49:15.870374Z"},"id":"1BfyvVAyePM5","outputId":"940bc14a-a64f-4f97-e360-7997f41f0205","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# state = torch.load(PATH2PROJECT / 'trainreface' / 'temp-best-checkout-resnet101.pth', map_location=lambda storage, loc: storage)\nstate = torch.load('best-checkout-finetune.pth', map_location=lambda storage, loc: storage)\n# state = torch.load('best-checkout.pth', map_location=lambda storage, loc: storage)\nbest_val_loss = state['best_val_loss']\nclassifier.load_state_dict(state['state_dict'])\n# classifier.to(device)","metadata":{"execution":{"iopub.status.busy":"2021-11-03T01:49:18.98738Z","iopub.execute_input":"2021-11-03T01:49:18.988513Z","iopub.status.idle":"2021-11-03T01:49:19.289744Z","shell.execute_reply.started":"2021-11-03T01:49:18.988466Z","shell.execute_reply":"2021-11-03T01:49:19.288799Z"},"id":"4bhCZv7XePM5","outputId":"fe9c15c9-eb3f-41d7-c815-a7c58103edbb","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from IPython.display import FileLink\n# FileLink('best-checkout.pth')","metadata":{"execution":{"iopub.status.busy":"2021-11-03T01:46:47.660831Z","iopub.execute_input":"2021-11-03T01:46:47.661801Z","iopub.status.idle":"2021-11-03T01:46:47.668252Z","shell.execute_reply.started":"2021-11-03T01:46:47.661765Z","shell.execute_reply":"2021-11-03T01:46:47.667216Z"},"id":"o0FHgsfAePM6","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# visualize_results(\n#     losses=losses,\n#     val_losses=val_losses,\n#     loglosses=loglosses,\n#     val_loglosses=val_loglosses,\n#     f1_scores=f1_scores,\n#     val_f1_scores=val_f1_scores\n# )","metadata":{"id":"d-dM-CHI0--h","execution":{"iopub.status.busy":"2021-11-02T22:56:14.594602Z","iopub.status.idle":"2021-11-02T22:56:14.595277Z","shell.execute_reply.started":"2021-11-02T22:56:14.59496Z","shell.execute_reply":"2021-11-02T22:56:14.594992Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference and submission","metadata":{"id":"Dec_5-LmePM8"}},{"cell_type":"code","source":"def inference(classifier, test_dataloader):\n    classifier.eval()\n    \n    all_preds = []\n    all_labels = []\n    with torch.no_grad():\n        for _, batch in enumerate(tqdm(test_dataloader, total=len(test_dataloader))):\n            # Make prediction.\n            y_pred = classifier(batch['faces'].to(device))\n\n            all_preds.extend(y_pred.squeeze(dim=-1).detach().cpu().numpy().tolist())\n            all_labels.extend(batch['label'].squeeze(dim=-1).numpy().tolist())\n    return all_preds, all_labels","metadata":{"execution":{"iopub.status.busy":"2021-11-03T01:49:22.18681Z","iopub.execute_input":"2021-11-03T01:49:22.187529Z","iopub.status.idle":"2021-11-03T01:49:22.203943Z","shell.execute_reply.started":"2021-11-03T01:49:22.187479Z","shell.execute_reply":"2021-11-03T01:49:22.202768Z"},"id":"eTKe6EZ5ePM9","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_prediction, _ = inference(classifier, test_dataloader)\nlen(test_prediction)","metadata":{"execution":{"iopub.status.busy":"2021-11-03T01:49:28.640029Z","iopub.execute_input":"2021-11-03T01:49:28.640329Z","iopub.status.idle":"2021-11-03T01:55:04.030191Z","shell.execute_reply.started":"2021-11-03T01:49:28.640299Z","shell.execute_reply":"2021-11-03T01:55:04.029175Z"},"id":"RX1m3zHgePM9","outputId":"f346b29e-00ac-4a05-d394-b52db19ab210","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_prediction, val_labels = inference(classifier, val_dataloader)\nlen(val_prediction)","metadata":{"execution":{"iopub.status.busy":"2021-11-03T01:55:04.032903Z","iopub.execute_input":"2021-11-03T01:55:04.033475Z","iopub.status.idle":"2021-11-03T01:56:20.98675Z","shell.execute_reply.started":"2021-11-03T01:55:04.033433Z","shell.execute_reply":"2021-11-03T01:56:20.985771Z"},"id":"JIHCf8SSePM9","outputId":"349463e1-1273-4ff7-af97-727c634cd9c7","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_test['score'] = test_prediction","metadata":{"execution":{"iopub.status.busy":"2021-11-03T01:56:20.98892Z","iopub.execute_input":"2021-11-03T01:56:20.989235Z","iopub.status.idle":"2021-11-03T01:56:21.00153Z","shell.execute_reply.started":"2021-11-03T01:56:20.989191Z","shell.execute_reply":"2021-11-03T01:56:21.000172Z"},"id":"aOZUTx4OePM-","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"thresholds = np.linspace(0, 1, len(np.unique(val_prediction)))\nf1_scores = [f1_score(val_labels, (np.array(val_prediction) > t).astype(np.uint8), average='micro') for t in tqdm(thresholds)]\nt_best = thresholds[np.argmax(f1_scores)]\nprint('Best threshold: ', t_best)\nprint('Best F1-Score: ', np.max(f1_scores))","metadata":{"execution":{"iopub.status.busy":"2021-11-03T01:56:21.004847Z","iopub.execute_input":"2021-11-03T01:56:21.00586Z","iopub.status.idle":"2021-11-03T01:57:43.768921Z","shell.execute_reply.started":"2021-11-03T01:56:21.005805Z","shell.execute_reply":"2021-11-03T01:57:43.76778Z"},"id":"43U9CE3DePM_","outputId":"a1ad1cf1-4ae3-46c7-9845-450a58e52e64","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from IPython.display import FileLink","metadata":{"execution":{"iopub.status.busy":"2021-11-03T01:57:43.771062Z","iopub.execute_input":"2021-11-03T01:57:43.771793Z","iopub.status.idle":"2021-11-03T01:57:43.777523Z","shell.execute_reply.started":"2021-11-03T01:57:43.771744Z","shell.execute_reply":"2021-11-03T01:57:43.776486Z"},"id":"UY2gwe9KePM_","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_test['label'] = (X_test['score'] > t_best).astype(int)\nsubmission_result = submission[['filename', 'path']].merge(X_test[['path', 'label']], on='path', how='left').fillna(0)[['filename', 'label']]\nsubmission_result['label'] = submission_result['label'].astype(int)\n\nassert submission_result.shape[0] == submission.shape[0]\n\nsubmission_result.to_csv('submission_result_th.csv', index=False)\n\nFileLink('submission_result_th.csv')","metadata":{"execution":{"iopub.status.busy":"2021-11-03T01:57:43.779449Z","iopub.execute_input":"2021-11-03T01:57:43.780183Z","iopub.status.idle":"2021-11-03T01:57:43.890414Z","shell.execute_reply.started":"2021-11-03T01:57:43.780137Z","shell.execute_reply":"2021-11-03T01:57:43.889389Z"},"id":"W9gn-zrlePM_","outputId":"fc4dd74d-a15a-41a4-e29d-1f41d48d5091","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"X_test['label'] = (X_test['score'] > 0.5).astype(int)\nsubmission_result = submission[['filename', 'path']].merge(X_test[['path', 'label']], on='path', how='left').fillna(0)[['filename', 'label']]\nsubmission_result['label'] = submission_result['label'].astype(int)\n\nassert submission_result.shape[0] == submission.shape[0]\n\nsubmission_result.to_csv('submission_result_05.csv', index=False)\n\nFileLink('submission_result_05.csv')","metadata":{"execution":{"iopub.status.busy":"2021-11-03T01:57:43.892066Z","iopub.execute_input":"2021-11-03T01:57:43.892655Z","iopub.status.idle":"2021-11-03T01:57:43.976855Z","shell.execute_reply.started":"2021-11-03T01:57:43.892581Z","shell.execute_reply":"2021-11-03T01:57:43.975693Z"},"id":"Lz62Z3DUePNA","outputId":"0196ada4-10a4-41d5-d767-153516c9ac37","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp submission_result_th.csv /content/drive/MyDrive/dl-creator-school/submission_result_th.csv","metadata":{"id":"XRqxgUwBRxow"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp submission_result_05.csv /content/drive/MyDrive/dl-creator-school/submission_result_05.csv","metadata":{"id":"YjDPunjlUXdB"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!cp best-checkout-finetune.pth /content/drive/MyDrive/dl-creator-school/best-checkout-finetune.pth","metadata":{"id":"EPbbWHPvWL1O"},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"id":"ZNAp5BcXWTXg"},"execution_count":null,"outputs":[]}]}