{"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":"# Import Libraries","metadata":{}},{"cell_type":"code","source":"import torch, torchvision\nfrom torchvision.transforms import ToTensor, Resize\nfrom torchvision import transforms\nimport os\nfrom PIL import Image\nfrom torch.utils.data import DataLoader\nimport torch.nn as nn\nfrom torch import optim\nfrom tqdm import tqdm\nimport copy\nimport numpy as np\nfrom torch.utils.data import Dataset\nimport requests\nfrom io import BytesIO","metadata":{"execution":{"iopub.status.busy":"2022-07-26T05:13:27.451670Z","iopub.execute_input":"2022-07-26T05:13:27.452351Z","iopub.status.idle":"2022-07-26T05:13:27.461891Z","shell.execute_reply.started":"2022-07-26T05:13:27.452297Z","shell.execute_reply":"2022-07-26T05:13:27.460647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Constants","metadata":{}},{"cell_type":"code","source":"TRAIN_PATH = \"./data/train\" # 10000 cats & 10000 dogs\nVAL_PATH = \"./data/val\" # 2500 cats & 2500 dogs extracted from train folder\nNUM_BATCH = 32\nEPOCHS = 5\nLEARNING_RATE = 1e-3\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"","metadata":{"execution":{"iopub.status.busy":"2022-07-26T05:14:14.253653Z","iopub.execute_input":"2022-07-26T05:14:14.254137Z","iopub.status.idle":"2022-07-26T05:14:14.261631Z","shell.execute_reply.started":"2022-07-26T05:14:14.254100Z","shell.execute_reply":"2022-07-26T05:14:14.260063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transformers","metadata":{}},{"cell_type":"code","source":"transform = transforms.Compose([\n    ToTensor(),\n    Resize((500,500))\n])","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Datset & DataLoader","metadata":{}},{"cell_type":"markdown","source":"#### Dataset Class","metadata":{}},{"cell_type":"code","source":"class CatDogDataset(Dataset):\n\n    def __init__(self, train_dir, transform = None):\n        \n        self.train_dir = train_dir\n        self.transform = transform\n        self.images = os.listdir(train_dir)\n        \n\n    def __len__(self):\n        return len(self.images)\n    \n    def __getitem__(self, index):\n        image_path = os.path.join(self.train_dir, self.images[index])\n        label = self.images[index].split(\".\")[0]\n\n        label = 0 if label == 'cat' else 1\n        \n        image = np.array(Image.open(image_path))\n        \n        if self.transform is not None:\n            image = self.transform(image)\n\n        return image, label","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Prepare data and dataloader","metadata":{}},{"cell_type":"code","source":"train_data = CatDogDataset(TRAIN_PATH, transform)\nval_data = CatDogDataset(VAL_PATH, transform)\n\n\ntrain_dl = DataLoader(train_data, batch_size=NUM_BATCH)\nval_dl = DataLoader(val_data, batch_size=NUM_BATCH)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Initialize ResNet18 ","metadata":{}},{"cell_type":"markdown","source":"#### Import Pretrained model","metadata":{}},{"cell_type":"code","source":"model = torchvision.models.resnet18(pretrained=True)\nprint(model)","metadata":{"execution":{"iopub.status.busy":"2022-07-26T05:14:34.478315Z","iopub.execute_input":"2022-07-26T05:14:34.479145Z","iopub.status.idle":"2022-07-26T05:14:34.742855Z","shell.execute_reply.started":"2022-07-26T05:14:34.479101Z","shell.execute_reply":"2022-07-26T05:14:34.742043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### freeze weights","metadata":{}},{"cell_type":"code","source":"for param in model.parameters():\n    param.requires_grad = False","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### finetune the last fully connected layer to prefered output","metadata":{}},{"cell_type":"code","source":"model.fc = nn.Sequential(*[\n    nn.Linear(in_features=512, out_features=2),\n    nn.Softmax(dim=1)\n])\nprint(model)","metadata":{"execution":{"iopub.status.busy":"2022-07-26T05:14:52.144815Z","iopub.execute_input":"2022-07-26T05:14:52.145324Z","iopub.status.idle":"2022-07-26T05:14:52.156339Z","shell.execute_reply.started":"2022-07-26T05:14:52.145291Z","shell.execute_reply":"2022-07-26T05:14:52.154817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Functions","metadata":{}},{"cell_type":"markdown","source":"#### Validation Function","metadata":{}},{"cell_type":"code","source":"def validate(model, data):\n\n    total = 0\n    correct = 0\n\n    for (images, labels) in data:\n        images = images.to(DEVICE)\n        x = model(images)\n        _, pred = torch.max(x, 1)\n        total += x.size(0)\n        correct += torch.sum(pred == labels)\n\n    return correct*100/total","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Train Function","metadata":{}},{"cell_type":"code","source":"def train(num_epoch = EPOCHS, lr = LEARNING_RATE, device = DEVICE):\n    accuracies = []\n    cnn = model.to(device)\n    cec = nn.CrossEntropyLoss()\n    optimizer = optim.Adam(cnn.parameters(), lr=lr)\n\n    max_accuracy = 0\n\n    for epoch in range(num_epoch):\n        for i, (images, labels) in tqdm(enumerate(train_dl)):\n            images = images.to(device)\n            labels = labels.to(device)\n            optimizer.zero_grad()\n            pred = cnn(images)\n            loss = cec(pred, labels)\n            loss.backward()\n            optimizer.step()\n        accuracy = float(validate(cnn,val_dl))\n        accuracies.append(accuracy)\n        if accuracy > max_accuracy:\n            best_model = copy.deepcopy(cnn)\n            max_accuracy = accuracy\n            print(\"saving best model with accuracy: \", accuracy)\n        print(\"Epoch: \", epoch+1, \"Accuracy: \", accuracy, \"%\")\n\n    return best_model\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train the model","metadata":{}},{"cell_type":"code","source":"resnet = train()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Save Best Model","metadata":{}},{"cell_type":"code","source":"torch.save(resnet.state_dict(), \"ResNet18_CatDog.pth\")","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"markdown","source":"#### Main Function","metadata":{}},{"cell_type":"code","source":"def inference(path, model, device=\"cpu\"):\n    try:\n        resp = requests.get(path, timeout=10)\n        print(\"request sent\")\n    except:\n        return False\n    \n    with torch.no_grad():\n        image = np.array(Image.open(BytesIO(resp.content)))\n        \n        image = transforms(image)\n        \n        image = image.unsqueeze(0)\n        pred = model(image.to(device))\n        return pred","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Run Inference","metadata":{}},{"cell_type":"code","source":"path = str(input(\"insert the image url: \"))\npred = inference(path, model)\nif torch.is_tensor(pred):\n    pred_idx = np.argmax(pred)\n\n    pred_label = \"cat\" if pred_idx == 0 else \"dog\"\n    \n    print(f\"Predicted: {pred_label}, Prob: {pred[0][pred_idx]*100}%\")\nelse:\n    print(\"can not get the url!!!\")","metadata":{},"execution_count":null,"outputs":[]}]}