{"cells":[{"metadata":{"_uuid":"3d026188-cab6-4d77-9073-85c864b5a5e5","_cell_guid":"0860bef1-6050-4883-9e2a-04faabce05e8","trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"67fef3af-9aab-4ceb-bdfc-6f8c6f5f31a3","_cell_guid":"d904ee65-9e5a-496f-ae6b-8f5c4e13f104","trusted":true},"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\n# import numpy as np # linear algebra\n# import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\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# import os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\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","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"aaba66a9-0e44-47f1-9a52-21282367b31e","_cell_guid":"5097fb7c-361e-4d9a-ad89-c2820c7267e3","trusted":true},"cell_type":"code","source":"import pytorch_lightning as pl\nimport torch\nimport os\nimport tensorflow as tf\nimport tensorflow.keras as keras\nimport matplotlib.pyplot as plt\nimport cv2\nimport numpy as np \nimport pandas as pd","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"e761acc9-396a-4727-93b2-bb55b95c9150","_cell_guid":"dd112737-99de-4e50-a338-8827aef67cfd","trusted":true},"cell_type":"code","source":"pl.__version__","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"e2c89235-65da-4ba8-ab05-92c16748bf0e","_cell_guid":"84992e0b-4e30-445e-94ec-cab081cb0173","trusted":true},"cell_type":"code","source":"torch.__version__","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"e56c6023-ef9b-4e11-81fe-00d14122e0a7","_cell_guid":"4dde904f-a012-4d64-879a-101d6b4ff4cf","trusted":true},"cell_type":"code","source":"import json\nfrom pprint import pprint\nJSON_PATH = '../input/cassava-leaf-disease-classification/label_num_to_disease_map.json'\nwith open(JSON_PATH) as f:\n    classes_dict = json.load(f)\npprint(classes_dict)\nprint(type(classes_dict))","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d8e4056e-711a-4973-af13-e155f17a5437","_cell_guid":"473f97ca-7153-4ec2-a659-29029740877a","trusted":true},"cell_type":"code","source":"TRAIN_CSV_PATH = '../input/cassava-leaf-disease-classification/train.csv'\ndf = pd.read_csv(TRAIN_CSV_PATH)\ndf","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"1a7cb2a0-1920-414d-a656-07e58ec06bc2","_cell_guid":"3436556e-05ae-4579-9487-d7788c0704de","trusted":true},"cell_type":"code","source":"df['disease_name'] = df['label'].apply(lambda x: classes_dict[str(x)])","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"f112ea76-53e7-4037-8d43-09f5e11d70ea","_cell_guid":"58bf5285-31c4-47e9-8a9c-6f24f1e26c15","trusted":true},"cell_type":"code","source":"def plot_class_histogram(df, title):\n    data = df['disease_name']\n    p = plt.hist(data)\n    plt.xticks(rotation='vertical')\n    plt.title(title)\n\n    plt.show()\n\nplot_class_histogram(df, 'full dataset classes')","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"233593ba-7b9e-4d5f-b725-375e9e299f43","_cell_guid":"5c5bd790-f088-4192-b129-b73d85b27c81","trusted":true},"cell_type":"code","source":"df['label_string'] =  df['label'].apply(lambda x: str(x))\nfrom sklearn.utils.class_weight import  compute_class_weight\nclass_weights = compute_class_weight('balanced',\n                                                 np.unique(df['label_string']),\n                                                 df['label_string'])\npprint(class_weights)\nclass_weights = {index: weights for index, weights in enumerate(class_weights)}\npprint(class_weights)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"cd3657ad-33e2-4571-9775-2fe43548de4a","_cell_guid":"58476c23-3af4-4d7e-a7fd-3227e289858a","trusted":true},"cell_type":"code","source":"from sklearn.model_selection import train_test_split\ndef split_dataset(df):\n    train_df, test_df = train_test_split(df, test_size=0.15, stratify=df['label_string'])\n    return train_df, test_df\n\ntrain_df, test_df = split_dataset(df)\ntrain_df.head()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"acfc8c96-2b3b-4596-aa51-4c7c47f1269b","_cell_guid":"bd04c4fe-fe8d-4322-84df-8116fff30e95","trusted":true},"cell_type":"code","source":"plot_class_histogram(train_df, 'train dataset classes')\nplot_class_histogram(test_df, 'test dataset classes')","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"27b5eecc-1fad-4e1c-9821-1ad274516df6","_cell_guid":"b05db8eb-d842-4d0d-b5d7-04e224932208","trusted":true},"cell_type":"code","source":"!pip install torchsummary","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"3953715f-7d77-4dea-96c7-8417acebf0b5","_cell_guid":"479cd7cd-74e0-4588-8bf4-36171a020b68","trusted":true},"cell_type":"code","source":"from albumentations.pytorch import ToTensorV2\nfrom pytorch_lightning.metrics.functional import f1, accuracy\nfrom skimage import io\nfrom sklearn import metrics\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import OneHotEncoder\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom sklearn.utils.class_weight import compute_class_weight, compute_sample_weight\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import DataLoader, random_split, WeightedRandomSampler\nfrom torch.utils.data import Dataset\nfrom torchsummary import summary\nfrom torchvision import transforms, models\nimport albumentations as A\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport os\nimport pandas as pd\nimport pytorch_lightning as pl\nimport seaborn as sns\nimport torch\nfrom pytorch_lightning.callbacks import ModelCheckpoint\nimport multiprocessing","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"4eee9b0b-dd52-4df4-94fd-c5d1e66be776","_cell_guid":"0c2e3deb-e36e-4d2f-b2bc-cf690f3f63ed","trusted":true},"cell_type":"code","source":"class CasavaDataset(Dataset):\n    def __init__(\n        self,\n        dataframe: pd.DataFrame,\n        img_dir: str,\n        encoder: OneHotEncoder,\n        transform=None,\n    ):\n\n        self.dataframe = dataframe  # pd.read_csv(csv_file)\n        self.img_dir = img_dir\n        self.transform = transform\n        self.encoder = encoder\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n\n        row = self.dataframe.iloc[idx]\n        img_id = row[\"image_id\"]\n        img_name = os.path.join(self.img_dir, img_id)\n        image = io.imread(img_name)\n\n        if self.transform:\n            # image = self.transform(image)\n            # old_image = image\n            image = self.transform(image=image)[\"image\"]\n\n        label = row[\"label\"]\n\n        encoded_label = self.encoder.transform(np.array([[label]]))\n        encoded_label = encoded_label.reshape(-1,)\n        sample = {\n            \"image\": image,\n            \"label\": encoded_label,\n        }  # , 'sample_weight':sample_weight}\n\n        return sample","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"473fdb16-c143-4d15-aeae-038a4db3c0d4","_cell_guid":"04b9ec3a-3604-472d-bf74-0129222f145f","trusted":true},"cell_type":"code","source":"class CasvaModel(pl.LightningModule):\n    def __init__(\n        self,\n        train_dataset,\n        val_dataset,\n        test_dataset,\n        class_weights,\n        training_samples_weights,\n        num_classes: int = 5,\n        image_dims=(3, 28, 28),\n        learning_rate=2e-4,\n        batch_size=128,\n    ):\n\n        super().__init__()\n        #self.device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n        self.training_samples_weights = training_samples_weights\n        self.batch_size = batch_size\n        self.class_weights = torch.from_numpy(class_weights)\n        self.train_dataset = train_dataset\n        self.val_dataset = val_dataset\n        self.test_dataset = test_dataset\n        self.learning_rate = learning_rate\n        self.num_classes = num_classes\n        self.img_dims = image_dims\n        self.data_loader_workers = multiprocessing.cpu_count()\n\n        self._prepare_model(self.num_classes, freeze_base_network=True)\n\n        self.criterion = nn.BCEWithLogitsLoss(pos_weight=self.class_weights).to(self.device)\n\n    def _prepare_model(self, num_classes, freeze_base_network):\n\n        resnet = models.resnext101_32x8d(pretrained=True)\n\n        if freeze_base_network:\n            for param in resnet.parameters():\n                param.requires_grad = False\n\n        num_ftrs = resnet.fc.in_features\n        resnet.fc = nn.Sequential(\n            nn.Dropout(p=0.2),\n            nn.Linear(in_features=num_ftrs, out_features=128),\n            nn.Linear(in_features=128, out_features=num_classes),\n        )\n\n        resnet =resnet.to(self.device)\n        if torch.cuda.is_available():\n            resnet.cuda()\n\n        self.model = resnet\n        print(self.img_dims)\n        summary(resnet, self.img_dims)\n\n    def forward(self, x):\n        result = self.model(x)\n        # result = F.sigmoid(d)\n        return result\n\n    def training_step(self, batch, batch_idx):\n        x = batch[\"image\"].to(self.device)\n        y = batch[\"label\"].to(self.device)\n\n        output = self(x)\n        loss = self.criterion(output, y.float())\n\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        x = batch[\"image\"].to(self.device)\n        y = batch[\"label\"].to(self.device)\n        output = self(x)\n        loss = self.criterion(output, y.float())\n\n        temp_logits = torch.sigmoid(output)\n        f1_value = f1(temp_logits, y, self.num_classes, average=\"macro\")\n\n        self.log(\"val_loss\", loss, prog_bar=True)\n        self.log(\"val_f1\", f1_value, prog_bar=True)\n        return loss\n\n    def test_step(self, batch, batch_idx):\n        return self.validation_step(batch, batch_idx)\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=self.learning_rate)\n        scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.1)\n        #scheduler2 = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=0, verbose=True)\n        return [optimizer], [scheduler]\n\n    def train_dataloader(self):\n        weights = self.training_samples_weights\n        sampler = WeightedRandomSampler(weights, len(weights))\n        return DataLoader(dataset=self.train_dataset,\n                          batch_size=self.batch_size,\n                          num_workers=self.data_loader_workers,\n                          sampler= sampler\n                          )\n\n    def val_dataloader(self):\n        return DataLoader(self.val_dataset, batch_size=self.batch_size,num_workers=self.data_loader_workers)\n\n    def test_dataloader(self):\n        return DataLoader(self.test_dataset, batch_size=self.batch_size,num_workers=self.data_loader_workers)\n","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"77820303-120d-44af-b9ba-d926e5a845e9","_cell_guid":"9405a189-4910-4be0-8f9b-295ef5d77b4e","trusted":true},"cell_type":"code","source":"def split_dataset(df, test_size: float = 0.15):\n    train_df, test_df = train_test_split(\n        df, test_size=test_size, random_state=42, stratify=df[\"label_string\"]\n    )\n    return train_df, test_df\n\n\ndef load_classes_from_json(json_dir: str):\n    with open(json_dir) as f:\n        classes_dict = json.load(f)\n    return classes_dict\n\n\ndef prepare_datasets(\n    csv_dir: str, json_dir: str, images_dir: str, train_transform, test_transform\n):\n    classes_dict = load_classes_from_json(json_dir)\n    df = pd.read_csv(csv_dir)\n    \n\n    df[\"label_string\"] = df[\"label\"].apply(lambda x: str(x))\n\n    class_weights = compute_class_weight(\n        \"balanced\", classes=np.unique(df[\"label_string\"]), y=df[\"label_string\"]\n    )\n\n\n\n    # df['sample_weight'] = samples_weights\n\n    train_df, test_df = split_dataset(df, test_size=0.1)\n    test_df, val_df = split_dataset(test_df, test_size=0.5)\n\n    encoder = OneHotEncoder()\n    encoder.sparse = False\n    encoder.fit(df[[\"label\"]])\n\n    training_samples_weights = compute_sample_weight(class_weight='balanced', y=train_df['label_string'])\n\n    train_dataset = CasavaDataset(train_df, images_dir, encoder, train_transform)\n    val_dataset = CasavaDataset(val_df, images_dir, encoder, test_transform)\n    test_dataset = CasavaDataset(test_df, images_dir, encoder, test_transform)\n\n    training_samples_weights = torch.DoubleTensor(training_samples_weights)\n\n    return (\n        train_dataset,\n        val_dataset,\n        test_dataset,\n        class_weights,\n        training_samples_weights,\n        classes_dict,\n        encoder,\n    )","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"7c1526c7-f0d6-4336-b029-1d5e558b75a8","_cell_guid":"db27654e-70ab-4194-ba06-206851a36e1a","trusted":true},"cell_type":"code","source":"def get_test_model_predictions(model, test_dataset, batch_size):\n\n    predictions, targets = [], []\n\n    loader = DataLoader(test_dataset, batch_size=batch_size)\n\n    for i_batch, sample_batched in enumerate(loader):\n        images = sample_batched[\"image\"].to(model.device)\n        true_labels = sample_batched[\"label\"].numpy().astype(np.uint8)\n\n        outputs = model(images)\n        results = torch.sigmoid(outputs).cpu().detach().numpy()\n        predictions.append(results)\n        targets.append(true_labels)\n\n    stacked_preds = np.concatenate(predictions, axis=0)\n    stacked_targets = np.concatenate(targets, axis=0)\n\n    return stacked_preds, stacked_targets\n\n\ndef test_with_metrics(preds, targets, class_dict, threshold=0.5):\n\n    rounded_preds = np.where(preds > threshold, 1, 0)\n\n    print(metrics.classification_report(rounded_preds, targets))\n\n    mcms = metrics.multilabel_confusion_matrix(targets, rounded_preds)\n\n    for class_idx, cf in enumerate(mcms):\n        class_name = class_dict[str(class_idx)]\n        print_confusion_matrix(cf, class_name)\n\n\ndef print_confusion_matrix(confusion_matrix, class_names, fontsize=14):\n\n    df_cm = pd.DataFrame(confusion_matrix,)\n\n    try:\n        heatmap = sns.heatmap(df_cm, annot=True, fmt=\"d\", cbar=False,)\n    except ValueError:\n        raise ValueError(\"Confusion matrix values must be integers.\")\n    heatmap.yaxis.set_ticklabels(\n        heatmap.yaxis.get_ticklabels(), rotation=0, ha=\"right\", fontsize=fontsize\n    )\n    heatmap.xaxis.set_ticklabels(\n        heatmap.xaxis.get_ticklabels(), rotation=0, ha=\"right\", fontsize=fontsize\n    )\n    plt.xlabel(\"True label\")\n    plt.ylabel(\"Predicted label\")\n    plt.title(\"Class - \" + class_names)\n    plt.show()\n","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"82832eb4-6981-4fd1-8bf2-fe05a22cc5aa","_cell_guid":"2c257a91-fe89-4113-9a80-4eca1c8dba7c","trusted":true},"cell_type":"code","source":"from pathlib import Path\nfrom skimage import io\nimport os\nimport torch\nimport numpy as np\nimport pandas as pd\n\ndef infer_casava(model,img_dir, test_transforms,decoder):\n\n\n\n    img_files = [str(x) for x in Path(img_dir).glob(\"*\")]\n\n    column1 = []\n    column2 = []\n    model.eval()\n\n\n    for img_name in img_files:\n        img_id = os.path.basename(img_name)\n        image = io.imread(img_name)\n        image = test_transforms(image=image)[\"image\"]\n        image = image.unsqueeze(dim=0)\n        image = image.to(model.device)\n\n        output = model(image)\n        results = torch.sigmoid(output).cpu().detach().numpy()\n        index = results.argmax()\n\n        hot_one = np.zeros((5,))\n        hot_one[index] = 1\n        label = decoder.inverse_transform(hot_one.reshape(1,-1)).item()\n\n        column1.append(img_id)\n        column2.append(label)\n\n    data = {'image_id':column1,\n            'label':column2\n            }\n\n    df = pd.DataFrame(data)\n\n    df.to_csv('submission.csv',index=False,)\n    print(df)\n\n","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"cac16be5-ac72-4979-8d43-de4ff370d123","_cell_guid":"5cb78e37-4146-4866-bc9b-ad33f64c77e1","trusted":true},"cell_type":"code","source":"def start_training(\n    csv_file: str,\n    json_file: str,\n    images_dir: str,\n    test_images_dir:str,\n    training_max_epochs: int,\n    finetuning_max_epochs: int,\n    img_size: int,\n    batch_size: int,\n    learning_rate: float,\n    use_gpus:bool,\n    viz_datasets: bool = True,\n):\n\n    IMG_SIZE_NO_CHANNEL = (img_size, img_size)  # todo pass it as argument\n    IMG_SIZE_CHANNEL = (3, *IMG_SIZE_NO_CHANNEL)\n\n    train_transform = A.Compose(\n        [\n\n            A.Resize(img_size, img_size),\n            A.ShiftScaleRotate(\n                shift_limit=0.05, scale_limit=0.05, rotate_limit=360, p=0.5\n            ),\n            A.VerticalFlip(),\n            A.HorizontalFlip(),\n            A.RandomGamma(),\n            A.CLAHE(),\n            A.GaussNoise(),\n            A.Cutout(),\n            A.Blur(),\n            A.RGBShift(r_shift_limit=20, g_shift_limit=20, b_shift_limit=20, p=0.5),\n            A.RandomBrightnessContrast(p=0.5),\n            A.Normalize(),#((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)),\n            ToTensorV2(),\n        ]\n    )\n\n    test_transform = A.Compose(\n        [\n            A.Resize(img_size, img_size),\n            A.Normalize(),#((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)),\n            ToTensorV2(),\n        ]\n    )\n\n    (\n        train_dataset,\n        val_dataset,\n        test_dataset,\n        class_weights,\n        training_samples_weights,\n        classes_dict,\n        encoder,\n    ) = prepare_datasets(\n        csv_file, json_file, images_dir, train_transform, test_transform\n    )\n\n    if viz_datasets:\n        for name, ds in [\n            (\"train\", train_dataset),\n            (\"validation\", val_dataset),\n            (\"test\", test_dataset),\n        ]:\n            viz_batch(ds, 10, title=name)\n\n    model = CasvaModel(\n        num_classes=len(classes_dict),\n        class_weights=class_weights,\n        training_samples_weights=training_samples_weights,\n        image_dims=IMG_SIZE_CHANNEL,\n        learning_rate=learning_rate,\n        train_dataset=train_dataset,\n        val_dataset=val_dataset,\n        test_dataset=test_dataset,\n        batch_size=batch_size,\n    )\n\n    checkpoint_callback = ModelCheckpoint(monitor='val_loss',\n                                          save_top_k=1,\n                                          verbose=True,\n                                          mode='min'\n                                          )\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n    trainer = pl.Trainer(\n        gpus=1 if use_gpus else 0,\n        max_epochs=training_max_epochs,\n        progress_bar_refresh_rate=20,\n        callbacks=[checkpoint_callback]\n    )\n    trainer.fit(model)\n    trainer.test(model)\n\n    #unfreeze model\n    model.unfreeze()\n    for param in model.model.parameters():\n        param.requires_grad = True\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    print(\"unfreezed all weights\")\n\n    fine_tuner = pl.Trainer(\n        gpus=1 if use_gpus else 0,\n        max_epochs=finetuning_max_epochs,\n        progress_bar_refresh_rate=20,\n        callbacks=[checkpoint_callback]\n    )\n\n    fine_tuner.fit(model)\n    fine_tuner.test(model)\n    \n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n    preds, targets = get_test_model_predictions(\n        model, test_dataset, batch_size=2\n    )\n    \n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n    test_with_metrics(preds, targets, classes_dict, threshold=0.5)\n    \n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\n    infer_casava(model,test_images_dir,test_transform,encoder)\n\n    print(\"done\")\n\n","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"ddc4dd6c-2b35-4d45-935e-b2e3494379d9","_cell_guid":"5329064b-bfc7-49ff-8e27-f1537e3a9d54","trusted":true},"cell_type":"code","source":"def viz_batch(ds, n_images, title):\n    a = 5\n    for i in range(n_images):\n        sample = ds[i]\n        x = sample[\"image\"].numpy()\n        y = sample[\"label\"]\n        img = x\n        label = y\n        plt.imshow(np.transpose(img, axes=(1, 2, 0)))\n        plt.title(f\"Set:{title}. Label:{label}\")\n        plt.tight_layout()\n        plt.show()","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"fc1a69be-eab3-4edc-a977-e294a70f51b3","_cell_guid":"c8073849-1386-4389-bde3-3bae37c3dc67","trusted":true},"cell_type":"markdown","source":"# Prepare creating submition file"},{"metadata":{"_uuid":"0267b142-979a-4ed1-9200-a99b002c3913","_cell_guid":"5033f82b-f981-4c59-b7c9-f35c912c47dd","trusted":true},"cell_type":"code","source":"    CSV_FILE = r\"../input/cassava-leaf-disease-classification/train.csv\"\n    JSON_FILE = r\"../input/cassava-leaf-disease-classification/label_num_to_disease_map.json\"\n    IMAGES_DIR = r\"../input/cassava-leaf-disease-classification/train_images\"\n    TEST_IMAGES_DIR = r\"../input/cassava-leaf-disease-classification/test_images\"\n    TRAINING_EPOCHS= 5\n    FINE_TUNING_EPOCHS= 10\n    IMG_SIZE = 512\n    LEARNING_RATE=2e-4\n    VIZ_DATA = False\n    BATCH_SIZE=8\n    USE_GPU = True\n\n    start_training(csv_file=CSV_FILE,\n                   json_file=JSON_FILE,\n                   images_dir=IMAGES_DIR,\n                   test_images_dir=TEST_IMAGES_DIR,\n                   training_max_epochs=TRAINING_EPOCHS,\n                   finetuning_max_epochs=FINE_TUNING_EPOCHS,\n                   img_size=IMG_SIZE,\n                   learning_rate=LEARNING_RATE,\n                   viz_datasets=VIZ_DATA,\n                   batch_size=BATCH_SIZE,\n                   use_gpus=USE_GPU,\n                   )","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}