{"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":"This is training notebook for a vanilla approach to the problem: take raw images and run them through a simple network. \nThe inference can be found here: https://www.kaggle.com/konradb/umnist-model-infer","metadata":{}},{"cell_type":"code","source":"!pip install tez\n!pip install timm","metadata":{"execution":{"iopub.status.busy":"2022-03-26T19:55:57.447006Z","iopub.execute_input":"2022-03-26T19:55:57.447512Z","iopub.status.idle":"2022-03-26T19:56:12.636651Z","shell.execute_reply.started":"2022-03-26T19:55:57.447451Z","shell.execute_reply":"2022-03-26T19:56:12.635640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport albumentations as A\nimport pandas as pd\nimport numpy as np\n\n\nfrom tez import Tez, TezConfig\nfrom tez.callbacks import EarlyStopping\nfrom tez.datasets import ImageDataset\n\n\nimport torch\nimport torch.nn as nn\nfrom torch.nn import functional as F\n\nfrom sklearn import metrics, model_selection, preprocessing\nimport timm\n\nfrom sklearn.model_selection import KFold\n\n# ignoring warnings\nimport warnings\nwarnings.simplefilter(\"ignore\")\n\nimport os, cv2, json\nfrom PIL import Image\n\nimport random","metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_kg_hide-input":true,"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","papermill":{"duration":5.732608,"end_time":"2021-12-20T22:53:39.01515","exception":false,"start_time":"2021-12-20T22:53:33.282542","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-03-26T19:56:12.645966Z","iopub.execute_input":"2022-03-26T19:56:12.646165Z","iopub.status.idle":"2022-03-26T19:56:14.738571Z","shell.execute_reply.started":"2022-03-26T19:56:12.646136Z","shell.execute_reply":"2022-03-26T19:56:14.737711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:    \n    # config\n    work_dir = '../input/ultra-mnist/'\n    img_folder = '../input/ultramnist-resized-256/'\n    batch_size = 16\n    epochs = 5\n    img_size = 256\n    seed = 42\n    target_size = 2\n    model = 'resnet50'\n    patience = 4 \n    nfolds = 5","metadata":{"papermill":{"duration":0.026232,"end_time":"2021-12-20T22:53:39.0593","exception":false,"start_time":"2021-12-20T22:53:39.033068","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-03-26T19:56:14.740511Z","iopub.execute_input":"2022-03-26T19:56:14.740782Z","iopub.status.idle":"2022-03-26T19:56:14.747161Z","shell.execute_reply.started":"2022-03-26T19:56:14.740745Z","shell.execute_reply":"2022-03-26T19:56:14.746423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed: int = 42) -> None:\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    \nseed_everything(CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2022-03-26T19:56:14.749739Z","iopub.execute_input":"2022-03-26T19:56:14.750091Z","iopub.status.idle":"2022-03-26T19:56:14.761759Z","shell.execute_reply.started":"2022-03-26T19:56:14.750055Z","shell.execute_reply":"2022-03-26T19:56:14.760984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Functions","metadata":{"papermill":{"duration":0.016944,"end_time":"2021-12-20T22:53:39.093163","exception":false,"start_time":"2021-12-20T22:53:39.076219","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class UModel(nn.Module):\n    def __init__(self, num_classes):\n        super().__init__()\n        self.model = timm.create_model(CFG.model, pretrained=True)\n        n_features = self.model.fc.in_features\n        self.model.fc = nn.Linear(n_features, num_classes)\n\n    def monitor_metrics(self, outputs, targets):\n        device = targets.get_device()\n        outputs = torch.argmax(outputs, dim=1).cpu().detach().numpy()\n        targets = targets.cpu().detach().numpy()\n        f1 = metrics.f1_score(targets, outputs, average=\"macro\")\n        accuracy = metrics.accuracy_score(targets, outputs)\n        return {\"acc\": torch.tensor(accuracy, device=device)}\n\n    def optimizer_scheduler(self):\n        opt = torch.optim.Adam(self.parameters(), lr=1e-3)\n        sch = torch.optim.lr_scheduler.ReduceLROnPlateau(\n            opt,\n            factor=0.5,\n            patience=2,\n            verbose=True,\n            mode=\"max\",\n            threshold=1e-4,\n        )\n        return opt, sch\n\n    def forward(self, image, targets=None):\n        outputs = self.model(image)\n        if targets is not None:\n            loss = nn.CrossEntropyLoss()(outputs, targets)\n            metrics = self.monitor_metrics(outputs, targets)\n            return outputs, loss, metrics\n        return outputs, 0, {}","metadata":{"execution":{"iopub.status.busy":"2022-03-26T19:56:14.763173Z","iopub.execute_input":"2022-03-26T19:56:14.763428Z","iopub.status.idle":"2022-03-26T19:56:14.776251Z","shell.execute_reply.started":"2022-03-26T19:56:14.763396Z","shell.execute_reply":"2022-03-26T19:56:14.775407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"aug = A.Compose([\n            A.Normalize(\n                mean=[0.5, 0.5, 0.5],\n                std=[0.5, 0.5, 0.5],\n                max_pixel_value=255.0, \n                p=1.0\n            ) ], p=1.)","metadata":{"execution":{"iopub.status.busy":"2022-03-26T19:56:14.777614Z","iopub.execute_input":"2022-03-26T19:56:14.777894Z","iopub.status.idle":"2022-03-26T19:56:14.787084Z","shell.execute_reply.started":"2022-03-26T19:56:14.777848Z","shell.execute_reply":"2022-03-26T19:56:14.786310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data\n","metadata":{"papermill":{"duration":0.017345,"end_time":"2021-12-20T22:53:39.269462","exception":false,"start_time":"2021-12-20T22:53:39.252117","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# training set\ndfx = pd.read_csv(CFG.work_dir + \"train.csv\")\ndfx['label'] = 0\ndfx.drop('digit_sum', axis = 1, inplace = True)\ndfx['path'] = CFG.img_folder + 'train_img/'\n\n# test set\ndfx_te = pd.read_csv(CFG.work_dir + 'sample_submission.csv')\ndfx_te['label'] = 1\ndfx_te.drop('digit_sum', axis = 1, inplace = True)\ndfx_te['path'] = CFG.img_folder + 'test_img/'\n\n# combine\ndfx = pd.concat([dfx, dfx_te], axis = 0)","metadata":{"execution":{"iopub.status.busy":"2022-03-26T19:56:14.788806Z","iopub.execute_input":"2022-03-26T19:56:14.789272Z","iopub.status.idle":"2022-03-26T19:56:14.844690Z","shell.execute_reply.started":"2022-03-26T19:56:14.789236Z","shell.execute_reply":"2022-03-26T19:56:14.843765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# split into folds\nkf = KFold(n_splits = 5, random_state = 42, shuffle = True)\nfold_id = np.zeros((len(dfx),1))\n\nfor (ii, (train_index, test_index)) in enumerate(kf.split(dfx)):\n    fold_id[test_index] = ii\n    \ndfx['fold'] = fold_id.astype(int)\n\n\ndfx.head()\n","metadata":{"execution":{"iopub.status.busy":"2022-03-26T19:56:14.847516Z","iopub.execute_input":"2022-03-26T19:56:14.847714Z","iopub.status.idle":"2022-03-26T19:56:14.871142Z","shell.execute_reply.started":"2022-03-26T19:56:14.847688Z","shell.execute_reply":"2022-03-26T19:56:14.870473Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"for fold in range(0,CFG.nfolds):\n    \n    # split\n    trn_idx = dfx[dfx['fold'] != fold].index\n    val_idx = dfx[dfx['fold'] == fold].index\n    df_train = dfx.loc[trn_idx].reset_index(drop=True)\n    df_valid = dfx.loc[val_idx].reset_index(drop=True)\n\n    train_image_paths = df_train.path + df_train.id + '.jpeg'\n    valid_image_paths = df_valid.path + df_valid.id + '.jpeg'\n    train_targets = df_train.label.values\n    valid_targets = df_valid.label.values\n    \n\n    # prepare datasets\n    train_dataset = ImageDataset(\n        image_paths=train_image_paths, targets=train_targets, \n        augmentations=aug)\n\n    valid_dataset = ImageDataset(\n        image_paths=valid_image_paths, targets=valid_targets,\n        augmentations=aug)\n    \n                \n    # fit model for this fold\n    model = UModel(num_classes = CFG.target_size) \n    model = Tez(model)\n    \n    es = EarlyStopping(\n        monitor=\"valid_acc\",\n        model_path = 'model_f' +str(fold) + '.bin',\n        patience = CFG.patience,\n        mode=\"max\",\n        save_weights_only=True,\n    )\n\n\n    config = TezConfig(\n        training_batch_size=CFG.batch_size,\n        validation_batch_size=CFG.batch_size,\n        epochs=CFG.epochs,\n        step_scheduler_after=\"epoch\",\n        step_scheduler_metric=\"valid_acc\",\n    )\n    model.fit(\n        train_dataset,\n        valid_dataset=valid_dataset,\n        config=config,\n        callbacks=[es],\n    )","metadata":{"execution":{"iopub.status.busy":"2022-03-26T19:58:40.997660Z","iopub.execute_input":"2022-03-26T19:58:40.998043Z","iopub.status.idle":"2022-03-26T20:16:06.385028Z","shell.execute_reply.started":"2022-03-26T19:58:40.998004Z","shell.execute_reply":"2022-03-26T20:16:06.383785Z"},"trusted":true},"execution_count":null,"outputs":[]}]}