{"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":"# Ink Detection DINO Self-Attention Approach\nhttps://github.com/facebookresearch/dino<br/>\nhttps://arxiv.org/abs/2104.14294","metadata":{"id":"SQJ1qfByPrMW"}},{"cell_type":"code","source":"import os\nimport cv2\nimport requests\nfrom io import BytesIO\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nimport torchvision\nfrom torchvision import transforms as pth_transforms\nimport numpy as np\nfrom PIL import Image","metadata":{"id":"GV8ILLyuQR6J","execution":{"iopub.status.busy":"2023-05-05T06:17:04.760361Z","iopub.execute_input":"2023-05-05T06:17:04.761128Z","iopub.status.idle":"2023-05-05T06:17:04.766375Z","shell.execute_reply.started":"2023-05-05T06:17:04.761084Z","shell.execute_reply":"2023-05-05T06:17:04.765409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tpaths=[]\nfor dirname, _, filenames in os.walk('/kaggle/input/vesuvius-challenge-ink-detection/train/1'):\n    for filename in filenames:\n        if filename[-4:]=='.tif':\n            tpaths+=[(os.path.join(dirname, filename))]\ntpaths=sorted(tpaths)\nprint(tpaths[0])\nprint(len(tpaths))","metadata":{"execution":{"iopub.status.busy":"2023-05-05T06:17:04.768031Z","iopub.execute_input":"2023-05-05T06:17:04.768366Z","iopub.status.idle":"2023-05-05T06:17:04.782045Z","shell.execute_reply.started":"2023-05-05T06:17:04.768337Z","shell.execute_reply":"2023-05-05T06:17:04.781449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"patch_size = 8\n\nmodel = torch.hub.load('facebookresearch/dino:main', 'dino_vits16')\n\ndevice = torch.device(\"cuda\") if torch.cuda.is_available() else torch.device(\"cpu\")\n\nfor p in model.parameters():\n        p.requires_grad = False\n        \nmodel.eval()\nmodel.to(device)","metadata":{"id":"ls4LK1rvO_8X","outputId":"1b177a73-7469-4897-9fdd-2703f0a20e42","execution":{"iopub.status.busy":"2023-05-05T06:17:04.783317Z","iopub.execute_input":"2023-05-05T06:17:04.783849Z","iopub.status.idle":"2023-05-05T06:17:05.222574Z","shell.execute_reply.started":"2023-05-05T06:17:04.783818Z","shell.execute_reply":"2023-05-05T06:17:05.221997Z"},"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## create_attentions","metadata":{"id":"BbrKTiUiXMm4"}},{"cell_type":"code","source":"def create_attentions(img0):\n\n    transform = pth_transforms.Compose([\n        pth_transforms.ToTensor(),\n        pth_transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n    ])\n    img = transform(img0)\n    # make the image divisible by the patch size\n    w, h = img.shape[1] - img.shape[1] % patch_size, img.shape[2] - img.shape[2] % patch_size\n    img = img[:, :w, :h].unsqueeze(0)\n    w_featmap = img.shape[-2] // patch_size\n    h_featmap = img.shape[-1] // patch_size\n    #attentions = model.forward_selfattention(img.to(device))\n    attentions = model.get_last_selfattention(img)   #img.cuda()\n    nh = attentions.shape[1] # number of head\n    # we keep only the output patch attention\n    attentions = attentions[0, :, 0, 1:].reshape(nh, -1)\n    # we keep only a certain percentage of the mass\n    val, idx = torch.sort(attentions)\n    val /= torch.sum(val, dim=1, keepdim=True)\n    cumval = torch.cumsum(val, dim=1)\n    threshold = 0.6 # We visualize masks obtained by thresholding the self-attention maps to keep xx% of the mass.\n    th_attn = cumval > (1 - threshold)\n    idx2 = torch.argsort(idx)\n    for head in range(nh):\n        th_attn[head] = th_attn[head][idx2[head]]\n    th_attn = th_attn.reshape(nh, w_featmap//2, h_featmap//2).float()\n    # interpolate\n    th_attn = nn.functional.interpolate(th_attn.unsqueeze(0), scale_factor=patch_size, mode=\"nearest\")[0].cpu().numpy()\n    attentions = attentions.reshape(nh, w_featmap//2, h_featmap//2)\n    attentions = nn.functional.interpolate(attentions.unsqueeze(0), scale_factor=patch_size, mode=\"nearest\")[0].cpu().numpy()\n    attentions_mean = np.mean(attentions, axis=0)\n\n    return attentions_mean","metadata":{"id":"uOqTknYmQZSO","outputId":"c226696c-0ac5-4846-f9a4-4f0682acbbad","execution":{"iopub.status.busy":"2023-05-05T06:17:05.230103Z","iopub.execute_input":"2023-05-05T06:17:05.230579Z","iopub.status.idle":"2023-05-05T06:17:05.241329Z","shell.execute_reply.started":"2023-05-05T06:17:05.230527Z","shell.execute_reply":"2023-05-05T06:17:05.240829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path=tpaths[4]\nimage=cv2.imread(path)\nimage=cv2.resize(image,dsize=(0,0),fx=0.1,fy=0.1)  \nplt.imshow(image)\nplt.axis('off') \nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-05T06:31:46.822563Z","iopub.execute_input":"2023-05-05T06:31:46.822849Z","iopub.status.idle":"2023-05-05T06:31:48.954633Z","shell.execute_reply.started":"2023-05-05T06:31:46.822809Z","shell.execute_reply":"2023-05-05T06:31:48.953812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image=cv2.imread(path)\nimage=cv2.resize(image,dsize=(0,0),fx=0.2,fy=0.2)  \n(h,w,c)=image.shape\nprint(h,h//3)\nprint(w,w//3)\nfig, axs = plt.subplots(3,3,figsize=(15,20))\nfor i in range(9):\n    r=i//3\n    c=i%3\n    imagei=image[(h//3)*r:(h//3)*(r+1),(w//3)*c:(w//3)*(c+1),:]\n    print(r,c,imagei.shape)\n\n    axs[r][c].imshow(imagei)\n    axs[r][c].axis('off') \nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-05T06:17:08.982721Z","iopub.execute_input":"2023-05-05T06:17:08.982937Z","iopub.status.idle":"2023-05-05T06:17:10.361759Z","shell.execute_reply.started":"2023-05-05T06:17:08.982888Z","shell.execute_reply":"2023-05-05T06:17:10.361034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image=cv2.imread(path)\nimage=cv2.resize(image,dsize=(0,0),fx=0.8,fy=0.8)  \n(h,w,c)=image.shape\nfig, axs = plt.subplots(3,3,figsize=(15,20))\nfor i in range(9):\n    r=i//3\n    c=i%3\n    imagei=image[(h//3)*r:(h//3)*(r+1),(w//3)*c:(w//3)*(c+1),:]\n    attentions_mean=create_attentions(imagei)\n    axs[r][c].imshow(attentions_mean)\n    axs[r][c].axis('off') \nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-05T06:17:38.791994Z","iopub.execute_input":"2023-05-05T06:17:38.792266Z","iopub.status.idle":"2023-05-05T06:18:58.921170Z","shell.execute_reply.started":"2023-05-05T06:17:38.792237Z","shell.execute_reply":"2023-05-05T06:18:58.919423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}