{"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":"code","source":"!pip install tez==0.2.0\n!pip install timm","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-03-24T06:33:20.999751Z","iopub.execute_input":"2022-03-24T06:33:21.000363Z","iopub.status.idle":"2022-03-24T06:33:38.706621Z","shell.execute_reply.started":"2022-03-24T06:33:21.000270Z","shell.execute_reply":"2022-03-24T06:33:38.705840Z"},"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\nimport tez\nfrom tez.datasets import ImageDataset\nfrom tez.callbacks import EarlyStopping\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 tqdm import tqdm\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-24T06:33:38.708699Z","iopub.execute_input":"2022-03-24T06:33:38.708992Z","iopub.status.idle":"2022-03-24T06:33:43.606487Z","shell.execute_reply.started":"2022-03-24T06:33:38.708956Z","shell.execute_reply":"2022-03-24T06:33:43.605717Z"},"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-512/'\n    model_folder = '../input/umnist-resnet152-model/'\n\n    batch_size = 16\n    epochs = 5\n    img_size = 512\n    seed = 42\n    target_size = 28\n    model = 'resnet152'\n    lr = 0.002\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-24T06:33:43.607805Z","iopub.execute_input":"2022-03-24T06:33:43.608083Z","iopub.status.idle":"2022-03-24T06:33:43.614170Z","shell.execute_reply.started":"2022-03-24T06:33:43.608051Z","shell.execute_reply":"2022-03-24T06:33:43.613524Z"},"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-24T06:33:43.616269Z","iopub.execute_input":"2022-03-24T06:33:43.617057Z","iopub.status.idle":"2022-03-24T06:33:43.629166Z","shell.execute_reply.started":"2022-03-24T06:33:43.617022Z","shell.execute_reply":"2022-03-24T06:33:43.627758Z"},"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(tez.Model):\n    def __init__(self, model_name, num_classes, learning_rate, n_train_steps):\n        super().__init__()\n        self.learning_rate = learning_rate\n        self.n_train_steps = n_train_steps\n        # Create Model\n        self.model = timm.create_model(model_name, \n                                       pretrained= False, \n                                       in_chans=3, \n                                       num_classes=num_classes)\n        self.step_scheduler_after = \"batch\"\n    \n    def monitor_metrics(self, outputs, targets):\n        if targets is None:\n            return {}\n        outputs = torch.argmax(outputs, dim=1).cpu().detach().numpy()\n        targets = targets.cpu().detach().numpy()\n        accuracy = metrics.accuracy_score(targets, outputs)\n        return {\"accuracy\": accuracy}\n    \n    \n    def fetch_optimizer(self):\n        opt = torch.optim.Adam(self.parameters(), lr=3e-4)\n        return opt\n    \n    def fetch_scheduler(self):\n        sch = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(self.optimizer, \n                                                                   T_0=10, \n                                                                   T_mult=1, \n                                                                   eta_min=1e-6, \n                                                                   last_epoch=-1)\n        return sch\n\n    def forward(self, image, targets=None):\n        x = self.model(image)\n        if targets is not None:\n            loss = nn.CrossEntropyLoss()(x, targets)\n            metrics = self.monitor_metrics(x, targets)\n            return x, loss, metrics\n        return x, 0, {}","metadata":{"execution":{"iopub.status.busy":"2022-03-24T06:33:43.635449Z","iopub.execute_input":"2022-03-24T06:33:43.636679Z","iopub.status.idle":"2022-03-24T06:33:43.658545Z","shell.execute_reply.started":"2022-03-24T06:33:43.636636Z","shell.execute_reply":"2022-03-24T06:33:43.656919Z"},"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-24T06:33:43.660374Z","iopub.execute_input":"2022-03-24T06:33:43.661121Z","iopub.status.idle":"2022-03-24T06:33:43.669357Z","shell.execute_reply.started":"2022-03-24T06:33:43.661085Z","shell.execute_reply":"2022-03-24T06:33:43.668656Z"},"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":"dfx = pd.read_csv(CFG.work_dir + \"train.csv\")\ndfx.rename(columns={\"id\": \"image_id\", \"digit_sum\": \"label\"}, inplace = True)   \ndfx['image_id'] = dfx['image_id'] + '.jpeg'\n\n\n# 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-24T06:33:43.670646Z","iopub.execute_input":"2022-03-24T06:33:43.675376Z","iopub.status.idle":"2022-03-24T06:33:43.777015Z","shell.execute_reply.started":"2022-03-24T06:33:43.674989Z","shell.execute_reply":"2022-03-24T06:33:43.776277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# prep test dataset\ndfx_te = pd.read_csv(CFG.work_dir + 'sample_submission.csv')\nte_orig_id = dfx_te['id'].copy()\ndfx_te['id'] = dfx_te['id'] + '.jpeg'\n\ntest_image_paths = [CFG.img_folder + 'test_img/' + x for x in dfx_te.id.values]\n\n# fake targets\ntest_targets = dfx_te.digit_sum.values\ntest_dataset = ImageDataset(\n    image_paths=test_image_paths,\n    targets=test_targets,\n    augmentations = aug,\n)","metadata":{"execution":{"iopub.status.busy":"2022-03-24T06:33:43.778294Z","iopub.execute_input":"2022-03-24T06:33:43.778969Z","iopub.status.idle":"2022-03-24T06:33:43.834100Z","shell.execute_reply.started":"2022-03-24T06:33:43.778934Z","shell.execute_reply":"2022-03-24T06:33:43.833143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# storage for oof and submission\nprval = np.zeros((dfx.shape[0], 28))\nprfull = np.zeros(( len(test_image_paths), 28))","metadata":{"execution":{"iopub.status.busy":"2022-03-24T06:33:43.839207Z","iopub.execute_input":"2022-03-24T06:33:43.839560Z","iopub.status.idle":"2022-03-24T06:33:43.848465Z","shell.execute_reply.started":"2022-03-24T06:33:43.839527Z","shell.execute_reply":"2022-03-24T06:33:43.847666Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# storage for oof and submission\nprval = np.zeros((dfx.shape[0], 28))\nprfull = np.zeros(( len(test_image_paths), 28))","metadata":{"execution":{"iopub.status.busy":"2022-03-24T06:33:43.853171Z","iopub.execute_input":"2022-03-24T06:33:43.853563Z","iopub.status.idle":"2022-03-24T06:33:43.861369Z","shell.execute_reply.started":"2022-03-24T06:33:43.853526Z","shell.execute_reply":"2022-03-24T06:33:43.860358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"'../input/resizedto512-image-ultramnist/test_img/olcqzjjmps.jpeg'","metadata":{}},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"for fold in range(CFG.nfolds):\n    \n    print('----------------------------------------')\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    image_path = CFG.img_folder + 'train_img/'\n    train_image_paths = [os.path.join(image_path, x) for x in df_train.image_id.values]\n    valid_image_paths = [os.path.join(image_path, x) for x in df_valid.image_id.values]\n    train_targets = df_train.label.values\n    valid_targets = df_valid.label.values\n \n    valid_dataset = ImageDataset(\n        image_paths=valid_image_paths, targets=valid_targets,\n        augmentations=aug)\n\n    \n    # instantiate and load model\n    n_train_steps = int(len(train_image_paths) / CFG.batch_size * CFG.epochs)\n    model = UModel(model_name = CFG.model, \n                   num_classes = CFG.target_size,\n                   learning_rate = CFG.lr, \n                   n_train_steps = n_train_steps) \n    model.load(CFG.model_folder + str(CFG.model) + 'model_es_s' +str(CFG.img_size)+'_f' + str(fold) + '.bin')\n    \n    print(fold)\n    \n    # produce predictions - oof \n    preds = model.predict(valid_dataset, batch_size= 128, n_jobs=-1) \n    temp_preds = None\n    for p in preds:\n        if temp_preds is None:\n            temp_preds = p\n        else:\n            temp_preds = np.vstack((temp_preds, p))      \n    prval[val_idx,:] = temp_preds\n    \n    \n    print(np.round(np.mean(np.argmax(temp_preds, axis=1) == dfx.label[val_idx]),4))\n    \n    # produce predictions - test data\n    preds = model.predict(test_dataset, batch_size= 128, n_jobs=-1) \n    temp_preds = None\n    for p in preds:\n        if temp_preds is None:\n            temp_preds = p\n        else:\n            temp_preds = np.vstack((temp_preds, p))      \n    break","metadata":{"execution":{"iopub.status.busy":"2022-03-24T06:33:43.862901Z","iopub.execute_input":"2022-03-24T06:33:43.863418Z","iopub.status.idle":"2022-03-24T06:45:38.850222Z","shell.execute_reply.started":"2022-03-24T06:33:43.863381Z","shell.execute_reply":"2022-03-24T06:45:38.848264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"# final_preds = np.argmax(prfull, axis=1)\nfinal_preds = temp_preds.argmax(axis = 1)\n\ndfx_te.digit_sum = final_preds\ndfx_te['id'] = te_orig_id\ndfx_te.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-03-24T06:45:38.851594Z","iopub.execute_input":"2022-03-24T06:45:38.851884Z","iopub.status.idle":"2022-03-24T06:45:38.953309Z","shell.execute_reply.started":"2022-03-24T06:45:38.851822Z","shell.execute_reply":"2022-03-24T06:45:38.952521Z"},"trusted":true},"execution_count":null,"outputs":[]}]}