{"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":"# Objectives\n1. Have above 70% accuracy\n2. Ideally, achieve it in 2 days max\n3. Ideally, stick to the initial Dataset","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-02-18T00:50:26.461143Z","iopub.execute_input":"2023-02-18T00:50:26.461520Z","iopub.status.idle":"2023-02-18T00:50:35.413361Z","shell.execute_reply.started":"2023-02-18T00:50:26.461488Z","shell.execute_reply":"2023-02-18T00:50:35.412218Z"}}},{"cell_type":"code","source":"!pip install tfrecord","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:03:13.475208Z","iopub.execute_input":"2023-02-19T22:03:13.476252Z","iopub.status.idle":"2023-02-19T22:03:30.282222Z","shell.execute_reply.started":"2023-02-19T22:03:13.476148Z","shell.execute_reply":"2023-02-19T22:03:30.280773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport tensorflow as tf\nimport numpy as np\nfrom sklearn.datasets import load_sample_image\nimport matplotlib.pyplot as plt\nfrom IPython.display import Image\nfrom tensorflow.train import BytesList, FloatList, Int64List\nfrom tensorflow.train import Example, Features, Feature\nimport tfrecord\nimport cv2\nimport glob\nimport torchvision\nimport os\nimport pandas as pd\nfrom torchvision.io import read_image\nimport torch\nfrom torch.utils.data import Dataset\nfrom torchvision import datasets\nfrom torchvision.transforms import ToTensor\nimport matplotlib.pyplot as plt\nimport copy\nimport torchvision.transforms as T\nimport random \nimport plotly as pt\nimport plotly.express as px\nimport os\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n\n\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\"\n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:03:30.284997Z","iopub.execute_input":"2023-02-19T22:03:30.285389Z","iopub.status.idle":"2023-02-19T22:03:41.627602Z","shell.execute_reply.started":"2023-02-19T22:03:30.285353Z","shell.execute_reply":"2023-02-19T22:03:41.626399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:03:41.628841Z","iopub.execute_input":"2023-02-19T22:03:41.629167Z","iopub.status.idle":"2023-02-19T22:03:41.635030Z","shell.execute_reply.started":"2023-02-19T22:03:41.629138Z","shell.execute_reply":"2023-02-19T22:03:41.633903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(device)\nif device == 'cuda':\n    cuda_id = torch.cuda.current_device()\n    print(\"CUDA Device ID: \", torch.cuda.current_device())\n    print(\"Name of the current CUDA Device: \", torch.cuda.get_device_name(cuda_id))","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:03:41.637502Z","iopub.execute_input":"2023-02-19T22:03:41.637879Z","iopub.status.idle":"2023-02-19T22:03:41.646895Z","shell.execute_reply.started":"2023-02-19T22:03:41.637845Z","shell.execute_reply":"2023-02-19T22:03:41.645715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dev or Demo mode","metadata":{}},{"cell_type":"code","source":"## Here we will skip the training process, and will use weight pretrained locally\n# mode = 'dev'\nmode = 'demo'","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:03:41.648770Z","iopub.execute_input":"2023-02-19T22:03:41.649178Z","iopub.status.idle":"2023-02-19T22:03:41.658221Z","shell.execute_reply.started":"2023-02-19T22:03:41.649144Z","shell.execute_reply":"2023-02-19T22:03:41.656987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Getting the Data","metadata":{}},{"cell_type":"code","source":"url_base = '/kaggle/input/tpu-getting-started/tfrecords-jpeg-'\nresolutions = ['192', '224', '331', '512']\nresolution = resolutions[1]\nfolders = ['train', 'val', 'test']\ndescription = {\n    \"id\": \"byte\",\n    \"image\": \"byte\"\n}\n\ndef transform_from_tsrecord_to_df(description, url_base, resolution, folder, is_class):\n    df = pd.DataFrame({'id': pd.Series(dtype='str'), 'class': pd.Series(dtype='int'), 'image': pd.Series(dtype='object')})\n    url = '{0}{1}x{1}/{2}/*'.format(url_base, resolution, folder)\n    if is_class:\n        description = description.copy()\n        description[\"class\"] = \"int\"\n\n    tf_files = glob.glob('{0}{1}x{1}/{2}/*'.format(url_base, resolution, folder))\n    for tf_file in tf_files:\n        loader = tfrecord.tfrecord_loader(tf_file, None, description)\n        for record in loader:\n            id_str = ''.join(map(chr, record['id']))\n            class_value = record['class'][0].item() if is_class else None\n\n            img_np = cv2.imdecode(record[\"image\"], cv2.IMREAD_COLOR)\n            img_np_rgb = cv2.cvtColor(img_np, cv2.COLOR_BGR2RGB)\n\n            df.loc[len(df.index)] = [id_str, class_value, img_np_rgb]\n    return df","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:03:41.660376Z","iopub.execute_input":"2023-02-19T22:03:41.660795Z","iopub.status.idle":"2023-02-19T22:03:41.674372Z","shell.execute_reply.started":"2023-02-19T22:03:41.660759Z","shell.execute_reply":"2023-02-19T22:03:41.673166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if mode == 'dev':\n    df_train = transform_from_tsrecord_to_df(description, url_base, resolution, folders[0], True)\n    print(df_train.dtypes)\n    print(df_train.shape)\n    print(type(df_train.loc[0]['image']))\n    plt.imshow(df_train.loc[0]['image'])","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:03:41.676310Z","iopub.execute_input":"2023-02-19T22:03:41.676857Z","iopub.status.idle":"2023-02-19T22:03:41.686846Z","shell.execute_reply.started":"2023-02-19T22:03:41.676780Z","shell.execute_reply":"2023-02-19T22:03:41.685514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if mode == 'dev':\n    plot_data = df_train.groupby(['class'])['id'].count().reset_index().rename(columns={\"id\": \"images_count\"})\n    fig = px.bar(plot_data, x='class', y='images_count')\n    fig.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:03:41.688560Z","iopub.execute_input":"2023-02-19T22:03:41.689018Z","iopub.status.idle":"2023-02-19T22:03:41.696895Z","shell.execute_reply.started":"2023-02-19T22:03:41.688983Z","shell.execute_reply":"2023-02-19T22:03:41.695740Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_validation = transform_from_tsrecord_to_df(description, url_base, resolution, folders[1], True)\nprint(df_validation.dtypes)\nprint(df_validation.shape)\nplt.imshow(df_validation.loc[0]['image'])","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:03:41.698457Z","iopub.execute_input":"2023-02-19T22:03:41.698857Z","iopub.status.idle":"2023-02-19T22:04:01.063717Z","shell.execute_reply.started":"2023-02-19T22:03:41.698815Z","shell.execute_reply":"2023-02-19T22:04:01.062803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_data = df_validation.groupby(['class'])['id'].count().reset_index().rename(columns={\"id\": \"images_count\"})\nfig = px.bar(plot_data, x='class', y='images_count')\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:01.067276Z","iopub.execute_input":"2023-02-19T22:04:01.067865Z","iopub.status.idle":"2023-02-19T22:04:02.037747Z","shell.execute_reply.started":"2023-02-19T22:04:01.067829Z","shell.execute_reply":"2023-02-19T22:04:02.036689Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if mode == 'dev':\n    df_test = transform_from_tsrecord_to_df(description, url_base, resolution, folders[2], False)\n    print(df_test.dtypes)\n    print(df_test.shape)\n    plt.imshow(df_test.loc[0]['image'])","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:02.039272Z","iopub.execute_input":"2023-02-19T22:04:02.039598Z","iopub.status.idle":"2023-02-19T22:04:02.045946Z","shell.execute_reply.started":"2023-02-19T22:04:02.039569Z","shell.execute_reply":"2023-02-19T22:04:02.044733Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"UNIQUE_LABEL_COUNT = len(df_validation['class'].unique())\nUNIQUE_LABEL_COUNT","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:02.047212Z","iopub.execute_input":"2023-02-19T22:04:02.047534Z","iopub.status.idle":"2023-02-19T22:04:02.065753Z","shell.execute_reply.started":"2023-02-19T22:04:02.047505Z","shell.execute_reply":"2023-02-19T22:04:02.064275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Augmentation mechanism","metadata":{}},{"cell_type":"code","source":"def image_transform(image):\n    \n    probability = 0.6\n    \n\n    grayscale_transformer=T.Compose([\n        T.ToPILImage(),\n        T.Grayscale(num_output_channels=3),\n    ])\n#     padding_transformer=T.Compose([\n#         T.ToPILImage(),\n#         T.Pad(padding=(int(random.randint(0, 20)), int(random.randint(0, 20))), fill=0.5), \n#     ])\n    rotation_transformer=T.Compose([\n        T.ToPILImage(),\n        T.RandomRotation(degrees=(0, 40)),\n    ])\n    crop_transformer=T.Compose([\n        T.ToPILImage(),\n        T.CenterCrop(180), \n        T.Resize(size=(224, 224)), \n    ])    \n    v_flipper=T.Compose([\n        T.ToPILImage(),\n        T.RandomVerticalFlip(p=probability),\n    ])\n    h_flipper=T.Compose([\n        T.ToPILImage(),\n        T.RandomHorizontalFlip(p=probability)\n    ])\n    perspective_transformer=T.Compose([\n        T.ToPILImage(),\n        T.RandomPerspective(distortion_scale=0.2, p=1.0)\n    ])\n    blur_transformer=T.Compose([\n        T.ToPILImage(),\n        T.GaussianBlur(kernel_size=(5, 9), sigma=(0.1, 1))\n    ])\n    posterize_transformer=T.Compose([\n        T.ToPILImage(),\n        T.RandomPosterize(bits=2)\n    ])\n    equalize_transformer=T.Compose([\n        T.ToPILImage(),\n        T.RandomEqualize()\n    ])\n    sharpness_transformer=T.Compose([\n        T.ToPILImage(),\n        T.RandomAdjustSharpness(sharpness_factor=2)\n    ])\n    autocontrast_transformer=T.Compose([\n        T.ToPILImage(),\n        T.RandomAutocontrast()\n    ])\n    solarize_transformer=T.Compose([\n        T.ToPILImage(),\n        T.RandomSolarize(threshold=192.0)\n    ])    \n    color_jitter_transformer=T.Compose([\n        T.ToPILImage(),\n        T.ColorJitter(brightness=.5, hue=.05)\n    ])   \n    invert_transformer=T.Compose([\n        T.ToPILImage(),\n        T.RandomInvert(p=0.02)\n    ])       \n#     erasing_transformer=T.Compose([\n#         T.ToPILImage(),\n#         T.RandomErasing(p=0.5, scale=(0.02, 0.33), ratio=(0.3, 3.3), value=0, inplace=False)\n#     ])       \n    \n    # Apply\n    \n#     if random.random() < 0.05:\n#         image = np.array(grayscale_transformer(image))\n       \n    if random.random() < probability:    \n        padding_transformer=T.Compose([\n            T.ToPILImage(),\n            T.Pad(padding=(random.randint(0, 20), random.randint(0, 20))), \n            T.Resize(size=(224, 224)), \n        ])\n        image = np.array(padding_transformer(image))\n        \n    \n    if random.random() < probability:\n        image = np.array(rotation_transformer(image))\n\n#     image = np.array(v_flipper(image))\n\n    image = np.array(h_flipper(image))  \n    \n    if random.random() < probability:\n        image = np.array(perspective_transformer(image))  \n        \n    if random.random() < 0.5:    \n        crop_transformer=T.Compose([\n            T.ToPILImage(),\n            T.CenterCrop(random.randint(150, 200)), \n            T.Resize(size=(224, 224)), \n        ])           \n        image = np.array(crop_transformer(image))  \n        \n        \n    if random.random() < 0.2:\n        image = np.array(blur_transformer(image))      \n        \n    if random.random() < probability:\n        image = np.array(posterize_transformer(image))     \n\n    if random.random() < probability:\n        image = np.array(equalize_transformer(image))  \n        \n    if random.random() < probability:\n        image = np.array(sharpness_transformer(image))   \n        \n    if random.random() < probability:\n        image = np.array(autocontrast_transformer(image))     \n\n    if random.random() < probability:\n        image = np.array(solarize_transformer(image))       \n        \n    if random.random() < probability:\n        image = np.array(color_jitter_transformer(image))    \n        \n#     if random.random() < probability:\n#         image = erasing_transformer(image)\n        \n    if random.random() < probability:\n        image = np.array(invert_transformer(image))\n    \n\n    return image","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:02.068053Z","iopub.execute_input":"2023-02-19T22:04:02.068581Z","iopub.status.idle":"2023-02-19T22:04:02.090519Z","shell.execute_reply.started":"2023-02-19T22:04:02.068483Z","shell.execute_reply":"2023-02-19T22:04:02.089674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Datasets and Dataloaders","metadata":{}},{"cell_type":"code","source":"class FlowersDataset(Dataset):\n    def __init__(self, df, UNIQUE_LABEL_COUNT, transform=None):\n        self.df = df\n        self.transform = transform\n\n    def __len__(self):\n        return self.df.shape[0]\n\n    def __getitem__(self, idx):\n        img_id = self.df.iloc[idx, 0]\n        label = self.df.iloc[idx, 1]\n        image = self.df.iloc[idx, 2]\n        \n        if self.transform:\n            image = self.transform(image)\n\n        y = np.zeros(UNIQUE_LABEL_COUNT, dtype=np.float32)\n        y[label] = int(1)\n        return img_id, y, image.T","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:02.091830Z","iopub.execute_input":"2023-02-19T22:04:02.092194Z","iopub.status.idle":"2023-02-19T22:04:02.106755Z","shell.execute_reply.started":"2023-02-19T22:04:02.092163Z","shell.execute_reply":"2023-02-19T22:04:02.105823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if mode == 'dev':\n    train_dataset = FlowersDataset(df_train, UNIQUE_LABEL_COUNT, image_transform)\n    train_loader = torch.utils.data.DataLoader(\n        train_dataset,\n        shuffle=True,\n        batch_size=32,\n    )","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:02.108007Z","iopub.execute_input":"2023-02-19T22:04:02.108515Z","iopub.status.idle":"2023-02-19T22:04:02.121748Z","shell.execute_reply.started":"2023-02-19T22:04:02.108483Z","shell.execute_reply":"2023-02-19T22:04:02.120838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"validation_dataset = FlowersDataset(df_validation, UNIQUE_LABEL_COUNT)\nvalidation_loader = torch.utils.data.DataLoader(\n    validation_dataset,\n    shuffle=True,\n    batch_size=32,\n)","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:02.123004Z","iopub.execute_input":"2023-02-19T22:04:02.123514Z","iopub.status.idle":"2023-02-19T22:04:02.132048Z","shell.execute_reply.started":"2023-02-19T22:04:02.123484Z","shell.execute_reply":"2023-02-19T22:04:02.131133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if mode == 'dev':\n    test_dataset = FlowersDataset(df_test, UNIQUE_LABEL_COUNT)\n    test_loader = torch.utils.data.DataLoader(\n        test_dataset,\n        batch_size=32,\n    )","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:02.133760Z","iopub.execute_input":"2023-02-19T22:04:02.134365Z","iopub.status.idle":"2023-02-19T22:04:02.143611Z","shell.execute_reply.started":"2023-02-19T22:04:02.134330Z","shell.execute_reply":"2023-02-19T22:04:02.142714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Models","metadata":{}},{"cell_type":"code","source":"models = [\n    {'name': 'resnet50', 'model': torchvision.models.resnet50(pretrained=True)},\n    {'name': 'resnext50_32x4d', 'model': torchvision.models.resnext50_32x4d(pretrained=True)},\n]","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:02.145080Z","iopub.execute_input":"2023-02-19T22:04:02.145663Z","iopub.status.idle":"2023-02-19T22:04:29.893956Z","shell.execute_reply.started":"2023-02-19T22:04:02.145629Z","shell.execute_reply":"2023-02-19T22:04:29.892571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Set out features to the number of our unique labels","metadata":{}},{"cell_type":"code","source":"# set out features to the number of our unique labels\ndef set_out_features(model):\n    model.fc = torch.nn.Sequential(\n        torch.nn.Linear(\n            in_features=2048,\n            out_features=UNIQUE_LABEL_COUNT\n        ),\n        torch.nn.Sigmoid()\n    )","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:29.895854Z","iopub.execute_input":"2023-02-19T22:04:29.896328Z","iopub.status.idle":"2023-02-19T22:04:29.903582Z","shell.execute_reply.started":"2023-02-19T22:04:29.896282Z","shell.execute_reply":"2023-02-19T22:04:29.902392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Adding dropout","metadata":{}},{"cell_type":"code","source":"# adding dropout\ndef append_dropout(model, rate=0.15):\n    for name, module in model.named_children():\n        if len(list(module.children())) > 0:\n            append_dropout(module)\n        if isinstance(module, torch.nn.ReLU):\n            new = torch.nn.Sequential(module, torch.nn.Dropout(p=rate, inplace=False))\n            setattr(model, name, new)","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:29.905196Z","iopub.execute_input":"2023-02-19T22:04:29.905651Z","iopub.status.idle":"2023-02-19T22:04:29.915664Z","shell.execute_reply.started":"2023-02-19T22:04:29.905609Z","shell.execute_reply":"2023-02-19T22:04:29.914644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loading params","metadata":{}},{"cell_type":"code","source":"def load_params(obj):\n    PATH = '/kaggle/input/add-to-competition-petals-to-the-metal/model_{0}_best.pt'.format(obj['name'])\n    params = torch.load(PATH, map_location='cpu')\n    obj['model'].load_state_dict(params)","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:29.917126Z","iopub.execute_input":"2023-02-19T22:04:29.917827Z","iopub.status.idle":"2023-02-19T22:04:29.926449Z","shell.execute_reply.started":"2023-02-19T22:04:29.917782Z","shell.execute_reply":"2023-02-19T22:04:29.925551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Apply","metadata":{}},{"cell_type":"code","source":"# setting out feature\nfor obj in models:\n    set_out_features(obj['model'])","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:29.927876Z","iopub.execute_input":"2023-02-19T22:04:29.928494Z","iopub.status.idle":"2023-02-19T22:04:29.946272Z","shell.execute_reply.started":"2023-02-19T22:04:29.928458Z","shell.execute_reply":"2023-02-19T22:04:29.945132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # dropuot\n# for obj in models:\n#     append_dropout(obj['model'])","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:29.947826Z","iopub.execute_input":"2023-02-19T22:04:29.948941Z","iopub.status.idle":"2023-02-19T22:04:29.953649Z","shell.execute_reply.started":"2023-02-19T22:04:29.948906Z","shell.execute_reply":"2023-02-19T22:04:29.952471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Load params\n# for obj in models:\n#     load_params(obj)","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:29.956492Z","iopub.execute_input":"2023-02-19T22:04:29.956919Z","iopub.status.idle":"2023-02-19T22:04:29.964526Z","shell.execute_reply.started":"2023-02-19T22:04:29.956884Z","shell.execute_reply":"2023-02-19T22:04:29.963096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train function","metadata":{}},{"cell_type":"code","source":"def train_models(models, epochs, train_loader, validation_loader):\n    train_results = []\n    \n    for obj in models:\n        print('----------------- {0} -----------------'.format(obj['name']))\n        model = obj['model']\n\n        if torch.cuda.is_available(): \n            model.cuda()\n\n        optimizer = torch.optim.AdamW(model.parameters())\n        criterion = torch.nn.BCELoss()\n\n        best_accuracy = 0\n        best_wts = None\n        best_df_validation_metric = pd.DataFrame(data=None,columns=['id','y','yhat'])\n        epochs_stats = pd.DataFrame(data=None,columns=['type', 'epoch_number', 'incorrect', 'correct', 'loss'])\n\n        for epoch_index in range(epochs):\n            model.train()    \n            epoch_loss = 0\n            correct = 0\n            incorrect = 0\n            batch_index = 0\n            processed_items = 0\n\n            for img_ids, y, images in train_loader:\n                if torch.cuda.is_available(): \n                    images = images.cuda()\n                    y = y.cuda()\n\n                batch_index += 1\n                processed_items += y.shape[0]\n\n                # clean frads\n                optimizer.zero_grad()\n\n                # forward\n                z = model(images.float())\n\n                # calculate loss\n                loss = criterion(z, y)\n\n                # multiply loss by batch size\n                epoch_loss += loss.data.item() * y.shape[0]\n\n                # backward\n                loss.backward()\n\n                # optimizer\n                optimizer.step()\n\n                # accuracy\n                _, labelhat = torch.max(z, 1)\n                _, label = torch.max(y, 1)\n                batch_correct = (labelhat == label).sum().item()\n                correct += batch_correct\n                incorrect += (y.shape[0] - batch_correct)\n                accuracy = correct / (incorrect + correct)\n\n                if processed_items%320 == 0:\n                    print('epoch {0}, items: {1},  loss: {2}, incorrect: {3} correct: {4}, accuracy: {5}'.format(epoch_index, processed_items, loss, incorrect, correct, accuracy))\n\n            print('ended epoch {0}, items: {1},  epoch loss: {2}, incorrect: {3} correct: {4}, accuracy: {5}'.format(epoch_index, processed_items, epoch_loss, incorrect, correct, accuracy))\n            epochs_stats.loc[len(epochs_stats.index)] = ['train', epoch_index, incorrect, correct, epoch_loss]\n\n            with torch.no_grad():\n                \n                model.eval()\n                correct = 0\n                incorrect = 0\n                df_validation_metric = pd.DataFrame(data=None,columns=['id','y','yhat'])\n\n                processed_items = 0\n                for img_ids, y, images in validation_loader:\n                    if torch.cuda.is_available(): \n                        images = images.cuda()\n                        y = y.cuda()\n\n                    processed_items += y.shape[0]\n\n                    # forward\n                    z = model(images.float())\n                    # accuracy\n                    _, labelhat = torch.max(z, 1)\n                    _, label = torch.max(y, 1)\n                    batch_correct = (labelhat == label).sum().item()\n                    correct += batch_correct\n                    incorrect += (y.shape[0] - batch_correct)\n                    accuracy = correct / (incorrect + correct)\n\n                    df_validation_metric = df_validation_metric.append(pd.DataFrame({'id': img_ids, 'y': label.cpu(), 'yhat': labelhat.cpu()}))\n\n                    if processed_items%320 == 0:\n                        print('validation, items: {0}, incorrect: {1} correct: {2}, accuracy: {3}'.format(processed_items, incorrect, correct, accuracy))\n\n                print('validation, total: {0} incorrect: {1} correct: {2} accuracy: {3}'.format(incorrect + correct, incorrect, correct, accuracy))\n                epochs_stats.loc[len(epochs_stats.index)] = ['validation', epoch_index, incorrect, correct, 0]\n\n                if accuracy > best_accuracy:\n                    print('best accuracy record: {0}'.format(accuracy))\n                    best_accuracy = accuracy\n                    best_df_validation_metric = df_validation_metric\n                    best_wts = copy.deepcopy(model.state_dict())\n        \n        # free gpu memory\n        model.cpu()\n        \n        # result\n        train_result = {\n            'model_name': obj['name'],\n            'current_wts': copy.deepcopy(model.state_dict()),\n            'best_wts': best_wts,\n            'best_accuracy': best_wts,\n            'best_df_validation_metric': best_df_validation_metric,\n            'epochs_stats': epochs_stats,\n        }        \n        train_results.append(train_result)\n\n    return train_results","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:29.966297Z","iopub.execute_input":"2023-02-19T22:04:29.966653Z","iopub.status.idle":"2023-02-19T22:04:29.992663Z","shell.execute_reply.started":"2023-02-19T22:04:29.966621Z","shell.execute_reply":"2023-02-19T22:04:29.991534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{}},{"cell_type":"code","source":"if mode == 'dev':\n    train_results = train_models(models, 10, train_loader, validation_loader)","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:29.994216Z","iopub.execute_input":"2023-02-19T22:04:29.995162Z","iopub.status.idle":"2023-02-19T22:04:30.008846Z","shell.execute_reply.started":"2023-02-19T22:04:29.995117Z","shell.execute_reply":"2023-02-19T22:04:30.007786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Metrics","metadata":{}},{"cell_type":"code","source":"def calc_metrics(df, y):\n    ROUND_DIGITS = 4\n    total = df.shape[0]\n    P     = df[df['y'] == y].shape[0]\n    N     = total - P\n    PP    = df[df['yhat'] == y].shape[0]\n    PN    = total - PP\n    assert(P + N == total)\n    \n    TP    = df[(df['y'] == y) & (df['yhat'] == y)].shape[0]\n    FN    = df[(df['y'] == y) & (df['yhat'] != y)].shape[0]\n    FP    = df[(df['y'] != y) & (df['yhat'] == y)].shape[0]\n    TN    = df[(df['y'] != y) & (df['yhat'] != y)].shape[0]\n    assert(TP + FN + FP + TN == total)\n    \n    prevalence = round(P / total, ROUND_DIGITS)\n    accuracy   = round((TP + TN) / total, ROUND_DIGITS)\n    \n    TPR = round(TP / P, ROUND_DIGITS)\n    FPR = round(FP / N, ROUND_DIGITS)\n     \n    presicion =  round(TP / (TP + FP), ROUND_DIGITS) if TP + FP > 0 else 0\n    recall = round(TP / (TP + FN), ROUND_DIGITS) if TP + FN > 0 else 0\n    \n    return [y, total, P, N, PP, PN, TP, FN, FP, TN, prevalence, accuracy, TPR, FPR, presicion, recall]\n\ndef format_df_metrics(df_results):\n    df_metrics = pd.DataFrame(data=None,columns=['y', 'total', 'P', 'N', 'PP', 'PN', 'TP', 'FN', 'FP', 'TN', 'prevalence', 'accuracy', 'TPR', 'FPR', 'presicion', 'recall'])\n    for y in sorted(df_results['y'].unique()):\n        df_metrics.loc[len(df_metrics.index)] = calc_metrics(df_results, y)\n\n    return df_metrics","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:30.012608Z","iopub.execute_input":"2023-02-19T22:04:30.012989Z","iopub.status.idle":"2023-02-19T22:04:30.026832Z","shell.execute_reply.started":"2023-02-19T22:04:30.012959Z","shell.execute_reply":"2023-02-19T22:04:30.025107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if mode == 'dev':\n    for train_result in train_results:\n        print('----------------- {0} -----------------'.format(train_result['model_name']))\n        print(train_result['epochs_stats'].head())\n        print(train_result['best_df_validation_metric'].head())\n        df_formatted_best_df_validation = format_df_metrics(train_result['best_df_validation_metric'])\n        print(df_formatted_best_df_validation)\n        print(df_formatted_best_df_validation[['presicion', 'recall']].mean())","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:30.033332Z","iopub.execute_input":"2023-02-19T22:04:30.033742Z","iopub.status.idle":"2023-02-19T22:04:30.041912Z","shell.execute_reply.started":"2023-02-19T22:04:30.033704Z","shell.execute_reply":"2023-02-19T22:04:30.040734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Save params","metadata":{}},{"cell_type":"code","source":"## saving params\n# if mode == 'dev':\n#     for train_result in train_results:\n#         torch.save(train_result['best_wts'], 'model_{0}_best.pt'.format(train_result['model_name']))","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:30.043738Z","iopub.execute_input":"2023-02-19T22:04:30.044203Z","iopub.status.idle":"2023-02-19T22:04:30.052196Z","shell.execute_reply.started":"2023-02-19T22:04:30.044159Z","shell.execute_reply":"2023-02-19T22:04:30.051318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Orchestre validation","metadata":{}},{"cell_type":"code","source":"def orchestrate_validation(models, data_loader, do_load_params):\n    \n    accuracies_per_model = {}\n    correct_ensemble = 0\n    incorrect_ensemble = 0\n    accuracy_ensemble = 0\n\n    df_validation_metric = pd.DataFrame(data=None,columns=['id','y','yhat'])\n    \n    processed_items = 0\n    \n    for obj in models:\n        # set to eval\n        obj['model'].eval()\n        \n        # loading params\n        if do_load_params:\n            load_params(obj)\n        \n        accuracies_per_model[obj['name']] = {'total': 0, 'accuracy': 0, 'correct': 0, 'incorrect': 0}\n        \n        if torch.cuda.is_available(): \n            obj['model'].cuda()        \n    \n    with torch.no_grad():    \n    \n        for img_ids, y, images in data_loader:\n            if torch.cuda.is_available(): \n                images = images.cuda()\n                y = y.cuda()\n\n            processed_items += y.shape[0]\n\n            # label\n            _, label = torch.max(y, 1)\n\n            z_ensemble = 0\n            # forward\n            for obj in models:\n                # forward\n                z = obj['model'](images.float())\n                _, labelhat = torch.max(z, 1)\n\n                batch_correct = (labelhat == label).sum().item()\n                accuracies_per_model[obj['name']]['correct'] += batch_correct\n                accuracies_per_model[obj['name']]['incorrect'] += (y.shape[0] - batch_correct)  \n                total = accuracies_per_model[obj['name']]['incorrect'] + accuracies_per_model[obj['name']]['correct']\n                accuracy = accuracies_per_model[obj['name']]['correct'] / total\n                accuracies_per_model[obj['name']]['accuracy'] = accuracy\n                accuracies_per_model[obj['name']]['total'] = total\n\n                z_ensemble += z / len(models)\n\n            # accuracy\n            _, labelhat_ensemble = torch.max(z_ensemble, 1)\n\n            batch_correct_ensemble = (labelhat_ensemble == label).sum().item()\n            correct_ensemble += batch_correct_ensemble\n            incorrect_ensemble += (y.shape[0] - batch_correct_ensemble)  \n            accuracy_ensemble = correct_ensemble / (incorrect_ensemble + correct_ensemble)\n\n            df_validation_metric = df_validation_metric.append(pd.DataFrame({'id': img_ids, 'y': label.cpu(), 'yhat': labelhat_ensemble.cpu()}))\n\n            if processed_items%320 == 0:\n                print('ensemble validation, items: {0}, accuracy: {1} correct: {2} incorrect: {3}'.format(processed_items, accuracy_ensemble, correct_ensemble, incorrect_ensemble))\n\n    print('---------------')\n    \n    for obj in models:\n        r = accuracies_per_model[obj['name']]\n        print('finish validation model name: {0} items: {1}, accuracy: {2} correct: {3} incorrect: {4}'.format(obj['name'],  r['total'], r['accuracy'], r['correct'], r['incorrect']))\n        \n    print('---------------')\n    \n    print('finish ensemble validation items: {0}, accuracy: {1} correct: {2} incorrect: {3}'.format(incorrect_ensemble + correct_ensemble, accuracy_ensemble, correct_ensemble, incorrect_ensemble))\n    \n    return df_validation_metric","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:30.053983Z","iopub.execute_input":"2023-02-19T22:04:30.054413Z","iopub.status.idle":"2023-02-19T22:04:30.073306Z","shell.execute_reply.started":"2023-02-19T22:04:30.054380Z","shell.execute_reply":"2023-02-19T22:04:30.071984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_validation_result = orchestrate_validation(models, validation_loader, True)","metadata":{"execution":{"iopub.status.busy":"2023-02-19T22:04:30.074871Z","iopub.execute_input":"2023-02-19T22:04:30.075354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Orchestre validation metrics","metadata":{}},{"cell_type":"code","source":"df_formatted_validation = format_df_metrics(df_validation_result)    \ndf_formatted_validation.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Precision","metadata":{}},{"cell_type":"code","source":"# precision\nfig = px.bar(df_formatted_validation, x='y', y='presicion')\nfig.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Recall","metadata":{}},{"cell_type":"code","source":"# recall\nfig = px.bar(df_formatted_validation, x='y', y='recall')\nfig.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Precision and Recall avarages","metadata":{}},{"cell_type":"code","source":"print(df_formatted_validation[['presicion', 'recall']].mean())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare submission","metadata":{}},{"cell_type":"code","source":"def orchestrate_test(models, data_loader):\n    \n    df_result = pd.DataFrame(data=None,columns=['id','label'])\n    \n    processed_items = 0\n    \n    for obj in models:\n        # set to eval\n        obj['model'].eval()\n        # loading params\n        load_params(obj)\n        \n        if torch.cuda.is_available(): \n            obj['model'].cuda()        \n            \n\n    with torch.no_grad():\n        \n        for img_ids, y, images in data_loader:\n            if torch.cuda.is_available(): \n                images = images.cuda()\n                y = y.cuda()\n\n            processed_items += y.shape[0]\n\n            # label\n            _, label = torch.max(y, 1)\n\n            z_ensemble = 0\n            # forward\n            for obj in models:\n                # forward\n                z = obj['model'](images.float())\n                _, labelhat = torch.max(z, 1)\n\n                z_ensemble += z / len(models)\n\n            # label\n            _, labelhat_ensemble = torch.max(z_ensemble, 1)\n\n            df_result = df_result.append(pd.DataFrame({'id': img_ids, 'label': labelhat_ensemble.cpu()}))\n\n            if processed_items%320 == 0:\n                print('ensemble test, items: {0}'.format(processed_items))\n\n    print('---------------')\n    \n    print('finish ensemble test items: {0}'.format(processed_items))\n    \n    df_result = df_result.set_index(['id'])\n    \n    return df_result","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if mode == 'dev':\n    df_result = orchestrate_test(models, test_loader)\nelse:\n    df_result = pd.read_csv('/kaggle/input/add-to-competition-petals-to-the-metal/submission.csv', index_col='id')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_result.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_result.to_csv('./submission.csv')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Final thoughts\nThe accuracy looks ok for submission. Spent 3 days on it. I was only using initial data. Something is done quite well, but there is a big room for improvement though.","metadata":{}}]}