{"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":"    import timm\n    import torch\n    from torch import nn\n    from torch.utils.data import Dataset, DataLoader\n    from sklearn.model_selection import train_test_split, StratifiedKFold\n    import pandas as pd\n    import ast\n    import matplotlib.pyplot as plt\n    import numpy as np","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.execute_input":"2023-09-24T08:55:28.174482Z","iopub.status.idle":"2023-09-24T08:55:28.18166Z","shell.execute_reply.started":"2023-09-24T08:55:28.174448Z","shell.execute_reply":"2023-09-24T08:55:28.18021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\ndevice","metadata":{"execution":{"iopub.status.busy":"2023-09-24T08:55:28.183427Z","iopub.execute_input":"2023-09-24T08:55:28.18388Z","iopub.status.idle":"2023-09-24T08:55:28.207724Z","shell.execute_reply.started":"2023-09-24T08:55:28.183827Z","shell.execute_reply":"2023-09-24T08:55:28.206449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pandas as pd\n\n# Путь к папке с файлами\npath = \"/kaggle/input/quick-draw-simplified-100files-shuffled/\"\n\n# Получение списка всех файлов в директории\nfiles = [f for f in os.listdir(path) if os.path.isfile(os.path.join(path, f))]","metadata":{"execution":{"iopub.status.busy":"2023-09-24T08:55:28.209287Z","iopub.execute_input":"2023-09-24T08:55:28.209925Z","iopub.status.idle":"2023-09-24T08:55:28.224044Z","shell.execute_reply.started":"2023-09-24T08:55:28.209843Z","shell.execute_reply":"2023-09-24T08:55:28.222843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_SIZE = 128\nimport cv2\ndef draw_cv2(raw_strokes, size=BASE_SIZE, thickness=6, time_color=True):\n    img = np.zeros((BASE_SIZE, BASE_SIZE), np.uint8) + 255\n    for t, stroke in enumerate(ast.literal_eval(raw_strokes)):\n        for i in range(len(stroke[0]) - 1):\n            _  = cv2.line(img, (stroke[0][i], stroke[1][i]), (stroke[0][i+1], stroke[1][i+1]), 0, thickness)\n    if size != BASE_SIZE:\n        return cv2.resize(img, (size, size))\n    else:\n        return img","metadata":{"execution":{"iopub.status.busy":"2023-09-24T09:00:45.816967Z","iopub.execute_input":"2023-09-24T09:00:45.817443Z","iopub.status.idle":"2023-09-24T09:00:45.827244Z","shell.execute_reply.started":"2023-09-24T09:00:45.817412Z","shell.execute_reply":"2023-09-24T09:00:45.825862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_drawings_to_image(file_path, save_path):\n    # Удостоверьтесь, что папка для сохранения существует\n    if not os.path.exists(save_path):\n        os.makedirs(save_path)\n\n    # Прочтите файл\n    data = pd.read_csv(file_path)\n    # Пройдитесь по каждой записи в файле\n    for idx, row in data.iterrows():\n        # Преобразуйте рисунок в изображение\n        img = draw_cv2(row['drawing'])\n        \n        # Создайте уникальное имя файла для сохранения\n        file_name = f\"{row['word']}-{idx}.png\"\n        \n        # Сохраните изображение\n        cv2.imwrite(os.path.join(save_path, file_name), img)\nfor file in files:\n    save_drawings_to_image(os.path.join(path, file), \"quickdraw_images\")","metadata":{"execution":{"iopub.status.busy":"2023-09-24T09:00:46.076245Z","iopub.execute_input":"2023-09-24T09:00:46.076651Z","iopub.status.idle":"2023-09-24T10:07:38.60646Z","shell.execute_reply.started":"2023-09-24T09:00:46.076624Z","shell.execute_reply":"2023-09-24T10:07:38.604643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!zip -r file.zip /kaggle/working/quickdraw_images","metadata":{"_kg_hide-output":true,"scrolled":true,"execution":{"iopub.status.busy":"2023-09-24T10:18:34.041759Z","iopub.execute_input":"2023-09-24T10:18:34.042427Z","iopub.status.idle":"2023-09-24T10:19:13.384935Z","shell.execute_reply.started":"2023-09-24T10:18:34.042384Z","shell.execute_reply":"2023-09-24T10:19:13.382382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Загрузка всех файлов в список DataFrames\ndfs = [pd.read_csv(path + file) for file in files]\n\n# Объединение всех DataFrames в один\ndata = pd.concat(dfs, ignore_index=True)","metadata":{"execution":{"iopub.status.busy":"2023-09-24T08:53:20.939237Z","iopub.execute_input":"2023-09-24T08:53:20.939736Z","iopub.status.idle":"2023-09-24T08:53:43.417564Z","shell.execute_reply.started":"2023-09-24T08:53:20.939691Z","shell.execute_reply":"2023-09-24T08:53:43.41657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data.shape","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#data = pd.read_csv(\"/kaggle/input/quick-draw-simplified-100files-shuffled/train_k0.csv\")\n#data","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_CLASSES = len(data['word'].unique())\nNUM_CLASSES","metadata":{"execution":{"iopub.status.busy":"2023-09-24T08:48:57.253377Z","iopub.execute_input":"2023-09-24T08:48:57.253716Z","iopub.status.idle":"2023-09-24T08:48:57.624024Z","shell.execute_reply.started":"2023-09-24T08:48:57.253686Z","shell.execute_reply":"2023-09-24T08:48:57.622597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"type(data['drawing'][0]), data['drawing'][0], type(data['word'][0]), data['word'][0]","metadata":{"execution":{"iopub.status.busy":"2023-09-24T08:48:57.626088Z","iopub.execute_input":"2023-09-24T08:48:57.626484Z","iopub.status.idle":"2023-09-24T08:48:57.636363Z","shell.execute_reply.started":"2023-09-24T08:48:57.62645Z","shell.execute_reply":"2023-09-24T08:48:57.6348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data.dtypes","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data['drawing'] = data['drawing'].apply(ast.literal_eval)","metadata":{"execution":{"iopub.status.busy":"2023-09-24T08:53:43.419559Z","iopub.execute_input":"2023-09-24T08:53:43.420422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"nrows, ncols = 3, 3\n\nfig, axes = plt.subplots(nrows=nrows, ncols=ncols, figsize=(6, 6))\n\nfor i, ax in enumerate(axes.ravel()):\n    if i >= len(data):  # Exit the loop if we have more subplots than data.\n        break\n    \n    drawing_list = data['drawing'][i]\n    \n    for segment in drawing_list:\n        x, y = segment\n        ax.plot(x, y)\n        \n    ax.set_title(data['word'][i])\n    ax.grid(False)\n    ax.axis('off')  # Hide axis for clarity\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class QuickDrawDataset(Dataset):\n    def __init__(self, data: pd.DataFrame) -> None:\n        self.data = data\n\n    def __len__(self) -> int:\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        img = draw_cv2(np.array(self.data['drawing'][idx], dtype=object))\n        img_tensor = torch.FloatTensor(img).unsqueeze(0)  # добавляем канал\n        img_tensor = img_tensor.repeat(3, 1, 1)  # дублируем канал, чтобы получить 3 канала\n        word = self.data['word_encoded'][idx]\n        return img_tensor, word\n\n\n\nclass QuickDrawModel(torch.nn.Module):\n    def __init__(self) -> None:\n        super().__init__()  # input size - any\n        self.backbone = timm.create_model(\n            \"efficientnet_b0\",\n            pretrained=True,\n            num_classes=NUM_CLASSES,\n        )\n        # no softmax\n        # (C, H, W)\n\n    def forward(self, x):\n        return self.backbone(x)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from typing import Tuple\ndef train_step(model: torch.nn.Module, \n               data_loader: torch.utils.data.DataLoader, \n               loss_fn: torch.nn.Module, \n               optimizer:torch.optim.Optimizer,\n               accuracy_fn,\n               device: torch.device = device) -> Tuple[float, float]:\n    train_loss, train_acc = 0, 0\n    for batch, (X, y) in enumerate(data_loader):\n        X, y = X.to(device), y.to(device)\n        model.train()\n        y_pred = model(X)\n        loss = loss_fn(y_pred, y)\n        train_loss += loss\n        train_acc += accuracy_fn(y_pred.argmax(dim = 1), y)\n        \n        optimizer.zero_grad()\n        \n        loss.backward()\n        optimizer.step()\n    \n    train_loss /= len(data_loader)\n    train_acc /= len(data_loader)\n    return train_loss, train_acc\n\ndef val_step(model: torch.nn.Module,\n               data_loader: torch.utils.data.DataLoader,\n               loss_fn: torch.nn.Module,\n               accuracy_fn,\n               device: torch.device = device) -> Tuple[float, float]:\n    test_loss, test_acc = 0, 0\n    model.eval()\n    with torch.inference_mode():\n        for X, y in data_loader:\n            X, y = X.to(device), y.to(device)\n            test_pred = model(X)\n            test_loss += loss_fn(test_pred, y)\n            test_acc += accuracy_fn(test_pred.argmax(dim = 1), y)\n        test_loss /= len(data_loader)\n        test_acc /= len(data_loader)\n        return test_loss, test_acc","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data, val_data = train_test_split(data, test_size = 0.2, random_state = 42)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = train_data.reset_index(drop=True)\nval_data = val_data.reset_index(drop=True)\n\ntrain_dataset = QuickDrawDataset(train_data)\nval_dataset = QuickDrawDataset(val_data)\n#test_dataset = QuickDrawDataset(test_data)\n\ntrain_loader = torch.utils.data.DataLoader(train_dataset, batch_size=16, shuffle=True)\nval_loader = torch.utils.data.DataLoader(val_dataset, batch_size=16)\n#test_loader = torch.utils.data.DataLoader(test_dataset)\n\nmodel = QuickDrawModel()\nmodel.to(device)\nprint()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_fn = nn.CrossEntropyLoss()\noptimizer = torch.optim.SGD(params = model.parameters(), lr = 0.1)\ndef accuracy_fn(y_pred, y_true):\n    correct = torch.eq(y_true, y_pred).sum().item()\n    acc = (correct/len(y_pred)) * 100\n    return acc","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"EPOCHS = 10\nfor epoch in range(EPOCHS):\n    train_loss, accuracy = train_step(model, train_loader, loss_fn, optimizer, accuracy_fn)\n    val_loss, val_accuracy = val_step(model, val_loader, loss_fn, accuracy_fn)\n    \n    print(f\"Epoch {epoch+1}/{EPOCHS}\")\n    print(f\"Train Loss: {train_loss:.4f}, Accuracy: {accuracy:.2f}%\")\n    print(f\"Validation Loss: {val_loss:.4f}, Validation Accuracy: {val_accuracy:.2f}%\")\n    print(\"-\" * 50)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 1 вариант\n# сохранение модели\ntorch.save(model, 'path_to_file.pt')\n# загрузка\n# model = torch.load('path_to_file.pt')\n# model = torch.load(\"/kaggle/input/quickdrawmodel\")\n# 2-й вариант\n# сохранение параметров\n#torch.save(model.state_dict(), 'path_to_parameters.pt')\n# загрузка\n# model = QuickDrawModel()\n# model.load_state_dict(torch.load('path_to_parameters.pt'))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\nimport torchvision.transforms as transforms\n\ndef load_image(image_path):\n    img = Image.open(image_path).convert('L')  # Преобразуем изображение в оттенки серого\n    img = img.resize((256, 256))  # Убедитесь, что размер соответствует входному размеру модели\n    \n    transform = transforms.Compose([\n        transforms.ToTensor(),\n        lambda x: x.repeat(3, 1, 1)  # Дублируем канал, чтобы получить 3 канала\n    ])\n    \n    img_tensor = transform(img)\n    return img_tensor\n\nimg = Image.open('/kaggle/input/quickdrawmodel/tree.png').convert('L')\nplt.imshow(img, cmap='gray')\n\nimg_tensor = load_image('/kaggle/input/quickdrawmodel/tree.png')\nmodel.eval()\n\n\nimg_tensor = img_tensor.unsqueeze(0).to(device)\nprediction = model(img_tensor)\npredicted_class = prediction.argmax(dim=1).item()\n\nprint(f\"Predicted class ID: {predicted_class}\")\nprint(f\"Predicted class name: {encoder.inverse_transform([predicted_class])[0]}\")\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Перерисовка csv to png (слишком тяжело)","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\n\n# Путь к папке с файлами\npath = \"quick-draw-simplified-100files-shuffled/\"\n\n# Получение списка всех файлов в директории\nfiles = [f for f in os.listdir(path) if os.path.isfile(os.path.join(path, f))]\nBASE_SIZE = 128\nimport cv2\ndef draw_cv2(raw_strokes, size=BASE_SIZE, thickness=6, time_color=True):\n    img = np.zeros((BASE_SIZE, BASE_SIZE), np.uint8) + 255\n    for t, stroke in enumerate(ast.literal_eval(raw_strokes)):\n        for i in range(len(stroke[0]) - 1):\n            _  = cv2.line(img, (stroke[0][i], stroke[1][i]), (stroke[0][i+1], stroke[1][i+1]), 0, thickness)\n    if size != BASE_SIZE:\n        return cv2.resize(img, (size, size))\n    else:\n        return img\n    def save_drawings_to_image(file_path, save_path):\n    # Удостоверьтесь, что папка для сохранения существует\n    if not os.path.exists(save_path):\n        os.makedirs(save_path)\n\n    # Прочтите файл\n    data = pd.read_csv(file_path)\n    # Пройдитесь по каждой записи в файле\n    for idx, row in data.iterrows():\n        # Преобразуйте рисунок в изображение\n        img = draw_cv2(row['drawing'])\n        \n        # Создайте уникальное имя файла для сохранения\n        file_name = f\"{row['word']}-{idx}.png\"\n        \n        # Сохраните изображение\n        cv2.imwrite(os.path.join(save_path, file_name), img)\nfor file in files:\n    save_drawings_to_image(os.path.join(path, file), \"quickdraw_images\")","metadata":{},"execution_count":null,"outputs":[]}]}