Team Ai
Modelpublic

MInference/v-niah-haystack

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
plot.ipynb6267 linesDownload Raw Back to root
1{2 "cells": [3  {4   "cell_type": "code",5   "execution_count": null,6   "metadata": {},7   "outputs": [],8   "source": [9    "from matplotlib import pyplot as plt\n",10    "import json\n",11    "\n",12    "with open(\"plots/lmms-lab/LLaVA-Video-7B-Qwen2_32_nextqa_0.01/sparse_ratio.json\", \"r\") as f:\n",13    "    data = json.load(f)\n",14    "\n",15    "# subplot per layer\n",16    "num_plots_per_row = 7\n",17    "\n",18    "fig, axs = plt.subplots(nrows=len(data) // num_plots_per_row, ncols=num_plots_per_row, figsize=(num_plots_per_row * 3, (len(data) // num_plots_per_row) * 2.5), sharex=True, sharey=False)\n",19    "for i, (layer, heads) in enumerate(data.items()):\n",20    "    # axs[i // 4, i % 4].plot(range(len(heads)), [head[\"overall_sr\"] for head in heads], label=\"Overall\")\n",21    "    axs[i // num_plots_per_row, i % num_plots_per_row].plot(range(len(heads)), [head[\"vision_sr\"] for head in heads], label=\"Vision\")\n",22    "    axs[i // num_plots_per_row, i % num_plots_per_row].plot(range(len(heads)), [head[\"language_sr\"] for head in heads], label=\"Language\")\n",23    "    axs[i // num_plots_per_row, i % num_plots_per_row].legend()\n",24    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_title(f\"Layer {i}\")\n",25    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xlabel(\"Head\")\n",26    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xticks(range(len(heads)))\n",27    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xticklabels([])\n",28    "# plt.savefig(\"plots/lmms-lab/LLaVA-Video-7B-Qwen2_32_nextqa_0.01/sparse_ratio.png\")\n",29    "fig.suptitle(\"Llava-Video-7B-Qwen2 - Sparse Ratio: $Num(critical\\_tokens)/Num(tokens)$\", fontsize=20, y=1.005)\n",30    "fig.tight_layout()\n",31    "plt.show()\n",32    "# plt.close()\n"33   ]34  },35  {36   "cell_type": "code",37   "execution_count": null,38   "metadata": {},39   "outputs": [],40   "source": [41    "from matplotlib import pyplot as plt\n",42    "import json\n",43    "\n",44    "with open(\"plots/Efficient-Large-Model/qwen2-7b-longvila-256f_32_nextqa_0.01/sparse_ratio.json\", \"r\") as f:\n",45    "    data = json.load(f)\n",46    "\n",47    "# subplot per layer\n",48    "num_plots_per_row = 7\n",49    "\n",50    "fig, axs = plt.subplots(nrows=len(data) // num_plots_per_row, ncols=num_plots_per_row, figsize=(num_plots_per_row * 3, (len(data) // num_plots_per_row) * 2.5), sharex=True, sharey=False)\n",51    "for i, (layer, heads) in enumerate(data.items()):\n",52    "    # axs[i // 4, i % 4].plot(range(len(heads)), [head[\"overall_sr\"] for head in heads], label=\"Overall\")\n",53    "    axs[i // num_plots_per_row, i % num_plots_per_row].plot(range(len(heads)), [head[\"vision_sr\"] for head in heads], label=\"Vision\")\n",54    "    axs[i // num_plots_per_row, i % num_plots_per_row].plot(range(len(heads)), [head[\"language_sr\"] for head in heads], label=\"Language\")\n",55    "    axs[i // num_plots_per_row, i % num_plots_per_row].legend()\n",56    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_title(f\"Layer {i}\")\n",57    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xlabel(\"Head\")\n",58    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xticks(range(len(heads)))\n",59    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xticklabels([])\n",60    "# plt.savefig(\"plots/lmms-lab/LLaVA-Video-7B-Qwen2_32_nextqa_0.01/sparse_ratio.png\")\n",61    "fig.suptitle(\"LongVILA-Qwen2 - Sparse Ratio: $Num(critical\\_tokens)/Num(tokens)$\", fontsize=20, y=1.005)\n",62    "fig.tight_layout()\n",63    "plt.show()\n",64    "# plt.close()\n"65   ]66  },67  {68   "cell_type": "code",69   "execution_count": null,70   "metadata": {},71   "outputs": [],72   "source": [73    "# compare language sparse ratio vs. vlm sparse ratio\n",74    "\n",75    "from matplotlib import pyplot as plt\n",76    "import json\n",77    "\n",78    "with open(\"plots/lmms-lab/LLaVA-Video-7B-Qwen2_32_nextqa_0.01/sparse_ratio.json\", \"r\") as f:\n",79    "    vlm_data = json.load(f)\n",80    "with open(\"plots/language_on_qwen2_7b/sparse_ratio.json\", \"r\") as f:\n",81    "    language_data = json.load(f)\n",82    "\n",83    "# subplot per layer\n",84    "num_plots_per_row = 7\n",85    "\n",86    "fig, axs = plt.subplots(nrows=len(vlm_data) // num_plots_per_row, ncols=num_plots_per_row, figsize=(num_plots_per_row * 3, (len(vlm_data) // num_plots_per_row) * 2.5), sharex=True, sharey=False)\n",87    "for i, (layer, heads) in enumerate(vlm_data.items()):\n",88    "    language_heads = language_data[layer]\n",89    "    axs[i // num_plots_per_row, i % num_plots_per_row].plot(range(len(heads)), [head[\"overall_sr\"] for head in language_heads], label=\"LLM\")\n",90    "    \n",91    "    axs[i // num_plots_per_row, i % num_plots_per_row].plot(range(len(heads)), [head[\"vision_sr\"] for head in heads], label=\"Vision\")\n",92    "    axs[i // num_plots_per_row, i % num_plots_per_row].plot(range(len(heads)), [head[\"language_sr\"] for head in heads], label=\"Language\")\n",93    "    axs[i // num_plots_per_row, i % num_plots_per_row].legend()\n",94    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_title(f\"Layer {i}\")\n",95    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xlabel(\"Head\")\n",96    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xticks(range(len(heads)))\n",97    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xticklabels([])\n",98    "fig.suptitle(\"Qwen2 (LLM) vs Llava-Video-Qwen2 (Vision & Language) - Sparse Ratio: $Num(critical\\_tokens)/Num(tokens)$\", fontsize=20, y=1.005)\n",99    "fig.tight_layout()\n",100    "plt.show()\n",101    "# plt.close()\n"102   ]103  },104  {105   "cell_type": "code",106   "execution_count": null,107   "metadata": {},108   "outputs": [],109   "source": [110    "# compare language sparse ratio vs. vlm sparse ratio\n",111    "\n",112    "from matplotlib import pyplot as plt\n",113    "import json\n",114    "\n",115    "with open(\"plots/lmms-lab/LLaVA-Video-7B-Qwen2_32_nextqa_0.01/block_coverage.json\", \"r\") as f:\n",116    "    vlm_data = json.load(f)\n",117    "with open(\"plots/language_on_qwen2_7b/block_coverage.json\", \"r\") as f:\n",118    "    language_data = json.load(f)\n",119    "\n",120    "# subplot per layer\n",121    "num_plots_per_row = 7\n",122    "\n",123    "fig, axs = plt.subplots(nrows=len(vlm_data) // num_plots_per_row, ncols=num_plots_per_row, figsize=(num_plots_per_row * 3, (len(vlm_data) // num_plots_per_row) * 2.5), sharex=True, sharey=False)\n",124    "for i, (layer, heads) in enumerate(vlm_data.items()):\n",125    "    language_heads = language_data[layer]\n",126    "    axs[i // num_plots_per_row, i % num_plots_per_row].plot(range(len(heads)), [head[\"blocks_ratio\"] for head in language_heads], label=\"LLM\")\n",127    "    \n",128    "    axs[i // num_plots_per_row, i % num_plots_per_row].plot(range(len(heads)), [head[\"blocks_ratio\"] for head in heads], label=\"Vision\")\n",129    "    axs[i // num_plots_per_row, i % num_plots_per_row].legend()\n",130    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_title(f\"Layer {i}\")\n",131    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xlabel(\"Head\")\n",132    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xticks(range(len(heads)))\n",133    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xticklabels([])\n",134    "fig.suptitle(\"Top-P: Num of blocks, tau=0.9\", fontsize=20, y=1.005)\n",135    "fig.tight_layout()\n",136    "plt.show()\n",137    "# plt.close()\n"138   ]139  },140  {141   "cell_type": "code",142   "execution_count": null,143   "metadata": {},144   "outputs": [],145   "source": [146    "# compare language sparse ratio vs. vlm sparse ratio\n",147    "\n",148    "from matplotlib import pyplot as plt\n",149    "import json\n",150    "\n",151    "with open(\"plots/lmms-lab/LLaVA-Video-7B-Qwen2_100_nextqa_0.01/block_coverage.json\", \"r\") as f:\n",152    "    vlm_data = json.load(f)\n",153    "with open(\"plots/language_on_qwen2_7b/block_coverage.json\", \"r\") as f:\n",154    "    language_data = json.load(f)\n",155    "\n",156    "# subplot per layer\n",157    "num_plots_per_row = 7\n",158    "\n",159    "fig, axs = plt.subplots(nrows=len(vlm_data) // num_plots_per_row, ncols=num_plots_per_row, figsize=(num_plots_per_row * 3, (len(vlm_data) // num_plots_per_row) * 2.5), sharex=True, sharey=False)\n",160    "for i, (layer, heads) in enumerate(vlm_data.items()):\n",161    "    language_heads = language_data[layer]\n",162    "    axs[i // num_plots_per_row, i % num_plots_per_row].plot(range(len(heads)), [head[\"blocks_ratio\"] for head in language_heads], label=\"LLM\")\n",163    "    \n",164    "    axs[i // num_plots_per_row, i % num_plots_per_row].plot(range(len(heads)), [head[\"blocks_ratio\"] for head in heads], label=\"Vision\")\n",165    "    axs[i // num_plots_per_row, i % num_plots_per_row].legend()\n",166    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_title(f\"Layer {i}\")\n",167    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xlabel(\"Head\")\n",168    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xticks(range(len(heads)))\n",169    "    axs[i // num_plots_per_row, i % num_plots_per_row].set_xticklabels([])\n",170    "fig.suptitle(\"Top-P: Num of blocks, tau=0.9\", fontsize=20, y=1.005)\n",171    "fig.tight_layout()\n",172    "plt.show()\n",173    "# plt.close()\n"174   ]175  },176  {177   "cell_type": "code",178   "execution_count": 2,179   "metadata": {},180   "outputs": [181    {182     "data": {183      "image/png": "iVBORw0KGgoAAAANSUhEUgAAAYUAAAGFCAYAAAASI+9IAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjAsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvlHJYcgAAAAlwSFlzAAAPYQAAD2EBqD+naQAABzBJREFUeJzt3T2u20YUgFEzUCVB0Aa0QC7JC/QWXE+qfDCM5IlWONT7OadSMXMx3YerhssYY3wDgG/fvv316gcA8H6IAgARBQAiCgBEFACIKAAQUQAgogBATlsPrsu1399//pjyGAAmOt8eHrEpABBRACCiAEBEAYCIAgARBQAiCgBEFACIKAAQUQAgogBARAGAiAIAEQUAIgoARBQAiCgAEFEAIKIAQEQBgIgCABEFAHJ65tJ6ue/9jnz/+WPabADeZlMAIKIAQJ76+2jvv3hm/h0FwHY2BQAiCgBEFACIKAAQUQAgogBARAGAiAIAEQUAIgoARBQAiCgAEFEAIKIAQEQBgIgCABEFACIKAEQUAMgyxhhbDq7LdfZbptv729IAH8r59vCITQGAiAIAOT1zae+/YdbL/ZDZALzNpgBARAGAiAIAEQUAIgoARBQAiCgAEFEAIKIAQEQBgIgCABEFACIKAEQUAIgoABBRACCiAEBEAYCIAgBZxhhjy8F1uc5+y4e297elAXZ3vj08YlMAIKIAQE7PXNr7r5L1cv/wswE+A5sCABEFACIKAEQUAIgoABBRACCiAEBEAYCIAgARBQAiCgBEFACIKAAQUQAgogBARAGAiAIAEQUAIgoAZBljjC0H1+U6+y38h72/LQ18UefbwyM2BQAiCgDk9Mylvf/OWC93s9+YDXAUmwIAEQUAIgoARBQAiCgAEFEAIKIAQEQBgIgCABEFACIKAEQUAIgoABBRACCiAEBEAYCIAgARBQAiCgBEFADIMsYYWw6uy3X2W3iR7z9/vPoJwBHOt4dHbAoARBQAyOmZS3v/3bBe7mYfOPv3+QD/sCkAEFEAIKIAQEQBgIgCABEFACIKAEQUAIgoABBRACCiAEBEAYCIAgARBQAiCgBEFACIKAAQUQAgogBAljHG2HJwXa6z38InNOP70sCTzreHR2wKAEQUAMjpmUt7/yWwXu5mHzh79vxfZwMfi00BgIgCABEFACIKAEQUAIgoABBRACCiAEBEAYCIAgARBQAiCgBEFACIKAAQUQAgogBARAGAiAIAEQUAsowxxpaD63Kd/Rb4IzO+XQ2f2vn28IhNAYCIAgA5PXNp77V9vdzNPnD27PlHzQb2Z1MAIKIAQEQBgIgCABEFACIKAEQUAIgoABBRACCiAEBEAYCIAgARBQAiCgBEFACIKAAQUQAgogBARAGALGOMseXgulxnvwXejRnfxYaXO98eHrEpABBRACCnZy7tvVqvl7vZB86ePf8zzIavyqYAQEQBgIgCABEFACIKAEQUAIgoABBRACCiAEBEAYCIAgARBQAiCgBEFACIKAAQUQAgogBARAGAiAIAWcYYY8vBdbnOfgt8CTO+uQ2bnG8Pj9gUAIgoAJDTM5f2Xn/Xy93sA2fPnm/227PhPbMpABBRACCiAEBEAYCIAgARBQAiCgBEFACIKAAQUQAgogBARAGAiAIAEQUAIgoARBQAiCgAEFEAIKIAQJYxxthycF2us98C/E8zvufNJ3K+PTxiUwAgogBATs9c2ntFXS93sw+cPXu+2a+bDf+XTQGAiAIAEQUAIgoARBQAiCgAEFEAIKIAQEQBgIgCABEFACIKAEQUAIgoABBRACCiAEBEAYCIAgARBQCyjDHGloPrcp39FuAdm/GtcA52vj08YlMAIKIAQE7PXNp7jVwvd7MPnD17vtmfczZfg00BgIgCABEFACIKAEQUAIgoABBRACCiAEBEAYCIAgARBQAiCgBEFACIKAAQUQAgogBARAGAiAIAEQUAsowxxpaD63Kd/Rbgi5rxHXL+xfn28IhNAYCIAgA5PXNp71VvvdzNPnD27Plmm/2ns3k/bAoARBQAiCgAEFEAIKIAQEQBgIgCABEFACIKAEQUAIgoABBRACCiAEBEAYCIAgARBQAiCgBEFACIKAAQUQAgyxhjbDm4LtfZbwGY4vvPH69+wvtwvj08YlMAIKIAQE7PXNp7FVsvd7MPnD17vtlmv3r27/PZzqYAQEQBgIgCABEFACIKAEQUAIgoABBRACCiAEBEAYCIAgARBQAiCgBEFACIKAAQUQAgogBARAGAiAIAWcYYY8vBdbnOfgvAhzPj+9LTnG8Pj9gUAIgoAJDTM5f2XpfWy93sA2fPnm+22a+ePXv+r7M/G5sCABEFACIKAEQUAIgoABBRACCiAEBEAYCIAgARBQAiCgBEFACIKAAQUQAgogBARAGAiAIAEQUAIgoAZBljjC0H1+U6+y0A/GL3b1efbw+P2BQAiCgAkNMzl/ZeadbL3ewDZ8+eb7bZr549e/5Rs1/BpgBARAGAiAIAEQUAIgoARBQAiCgAEFEAIKIAQEQBgIgCABEFACIKAEQUAIgoABBRACCiAEBEAYCIAgBZxhjj1Y8A4H2wKQAQUQAgogBARAGAiAIAEQUAIgoARBQAiCgAkL8Bq2La7A/v7IsAAAAASUVORK5CYII=",184      "text/plain": [185       "<Figure size 640x480 with 1 Axes>"186      ]187     },188     "metadata": {},189     "output_type": "display_data"190    }191   ],192   "source": [193    "import torch\n",194    "import seaborn as sns\n",195    "import matplotlib.pyplot as plt\n",196    "\n",197    "num_tokens = 128\n",198    "stride = 8\n",199    "\n",200    "def plot_mask(mask):\n",201    "    plt.figure(figsize=(8, 8), dpi=120)\n",202    "    sns.heatmap(mask.numpy(), cbar=False)\n",203    "    plt.axis('off')\n",204    "    plt.show()\n",205    "\n",206    "def plot_imshow(mask):\n",207    "    plt.imshow(\n",208    "        mask.numpy(),\n",209    "        # cmap=\"binary\",\n",210    "        # cmap=\"Blues\",\n",211    "        # cmap=\"viridis\",\n",212    "        cmap=\"Reds\",\n",213    "        interpolation=\"nearest\",\n",214    "        vmin=0,\n",215    "        vmax=1\n",216    "    )\n",217    "    plt.axis('off')\n",218    "    plt.show()\n",219    "\n",220    "# Creating the initial mask\n",221    "mask = torch.zeros((num_tokens, num_tokens), dtype=torch.int32)\n",222    "# for i in range(num_tokens):\n",223    "#     mask[i, i] = 1\n",224    "    # for j in range(0, i, stride):\n",225    "    #     mask[i, i - j] = 1\n",226    "\n",227    "# Adding diagonal elements with stride\n",228    "for i in range(0, num_tokens, stride):\n",229    "    mask[i, i] = 1\n",230    "    mask[i:, i] = 1\n",231    "    mask[i, :i] = 1\n",232    "\n",233    "# plot_mask(mask)\n",234    "plot_imshow(mask)"235   ]236  },237  {238   "cell_type": "code",239   "execution_count": 35,240   "metadata": {},241   "outputs": [242    {243     "data": {244      "text/plain": [245       "tensor([[  7,  15,  23,  31,  39,  47,  55,  63,  71,  79,  87,  95, 103, 111,\n",246       "         119, 127],\n",247       "        [  6,  14,  22,  30,  38,  46,  54,  62,  70,  78,  86,  94, 102, 110,\n",248       "         118, 126],\n",249       "        [  5,  13,  21,  29,  37,  45,  53,  61,  69,  77,  85,  93, 101, 109,\n",250       "         117, 125],\n",251       "        [  4,  12,  20,  28,  36,  44,  52,  60,  68,  76,  84,  92, 100, 108,\n",252       "         116, 124],\n",253       "        [  3,  11,  19,  27,  35,  43,  51,  59,  67,  75,  83,  91,  99, 107,\n",254       "         115, 123],\n",255       "        [  2,  10,  18,  26,  34,  42,  50,  58,  66,  74,  82,  90,  98, 106,\n",256       "         114, 122],\n",257       "        [  1,   9,  17,  25,  33,  41,  49,  57,  65,  73,  81,  89,  97, 105,\n",258       "         113, 121],\n",259       "        [  0,   8,  16,  24,  32,  40,  48,  56,  64,  72,  80,  88,  96, 104,\n",260       "         112, 120]])"261      ]262     },263     "execution_count": 35,264     "metadata": {},265     "output_type": "execute_result"266    }267   ],268   "source": [269    "torch.arange(num_tokens, dtype=torch.int64).reshape((-1, stride)).T.flip(0)\n"270   ]271  },272  {273   "cell_type": "code",274   "execution_count": 49,275   "metadata": {},276   "outputs": [277    {278     "data": {279      "image/png": "iVBORw0KGgoAAAANSUhEUgAAAYUAAAGFCAYAAAASI+9IAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjAsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvlHJYcgAAAAlwSFlzAAAPYQAAD2EBqD+naQAAB95JREFUeJzt2zFSxNgZRlF6yhFTFBvwAlmSF+gtTPwmgLomACOYFnp6OidSoIDsr++quY0xxgMAPDw8/HH0HwDAPBwFAOIoABBHAYA4CgDEUQAgjgIAcRQAyL+2vvhye+r5P3/9d5c/BoAdPT5/+YqlAEAcBQCyOR/t6eXPf/csTQEcx1IAII4CAJkiH+1JmgLYzlIAII4CAFk+H+1JmgJWYykAEEcBgMhHk5KmgCNYCgDEUrgoSwT4iKUAQBwFACIfcXfSFJyXpQBAHAUAIh9xKtIU7MtSACCOAgCRj+CNNAWWAgDvOAoARD6CXyBNcRaWAgBxFACIfAQnJ01xT5YCAHEUAIh8BHxKmroeSwGAOAoARD4CDiFNzclSACCOAgCRj4AlyVM/YykAEEsB4JtWXiGWAgBxFACIfAQwkaPTlKUAQBwFAHIbY4wtL77cnnpe7Ws7wCU8Pn/5iqUAQBwFADLFr4+O/toOwCtLAYA4CgBkiny0J2kKYDtLAYA4CgBk+Xy0J2kKWI2lAEAcBQAiH01KmgKOYCkAEEcBgMhHFyVPAR+xFACIowBA5CPuTpqC87IUAIilwKlYIbAvSwGAOAoARD6CN9IUWAoAvOMoABD5CH6BNMVZWAoAxFEAIPIRnJw0xT1ZCgDEUQAg8hHwKWnqeiwFAOIoABD5CDiENDUnSwGAOAoARD4CliRP/YylAEAcBQAiHwF808ppylIAII4CAJGPACZydJqyFADIbYwxtrz4cnvqebUPKwCX8Pj85SuWAgBxFADIFB+aj/6wAsArSwGAOAoAZIp8tCdpCmA7SwGAOAoAZPl8tCdpCliNpQBAHAUAIh9NSpoCjmApABBHAYDIRxclTwEfsRQAiKMAQOQj7k6agvOyFACIowBA5CNORZqCfVkKAMRRACDyEbyRpsBSAOAdRwGAyEfwC6QpzsJSACCWApycFcI9WQoAxFEAIPIR8Clp6nosBQDiKAAQ+Qg4hDQ1J0sBgDgKAEQ+ApYkT/2MpQBAHAUAIh8BfNPKacpSACCOAgCRjwAmcnSashQAiKMAQG5jjLHlxZfbU8+rfW0HuITH5y9fsRQAiKMAQKb49dHRX9sBeGUpABBHAYBMkY/2JE0BbGcpAJDll8KerBBgNZYCAHEUAIh8NClpCjiCpQBAHAUAIh9dlDwFfMRSACCOAgCRj7g7aQrOy1IAII4CAJGPOBVpCvZlKQAQRwGAyEfwRpoCSwGAdxwFACIfwS+QpjgLSwGAOAoARD6Ck5OmuCdLAYA4CgBEPgI+JU1dj6UAQCwF4BBWyJwsBQDiKAAQ+QhYkjz1M5YCAHEUAIh8BPBNK6cpSwGAOAoARD4CmMjRacpSACCOAgC5jTHGlhdfbk89r/a1HeASHp+/fMVSACCOAgCZ4tdHR39tB+CVpQBAHAUAMkU+2pM0BbCdpQBAHAUAsnw+2pM0BazGUgAgjgIAkY8mJU0BR7AUAIilcFGWCPARSwGAOAoARD7i7qQpOC9LAYA4CgBEPuJUpCnYl6UAQBwFACIfwRtpCiwFAN5xFACIfAS/QJriLCwFAOIoABD5CE5OmuKeLAUA4igAEPkI+JQ0dT2WAgBxFACIfAQcQpqak6UAQBwFACIfAUuSp37GUgAglgLAN628QiwFAOIoABD5CGAiR6cpSwGAOAoA5DbGGFtefLk99bza13aAS3h8/vIVSwGAOAoAZIpfHx39tR2AV5YCAHEUAMgU+WhP0hTAdpYCAHEUAMjy+WhP0hSwGksBgDgKAEQ+mpQ0BRzBUgAgjgIAkY8uSp4CPmIpABBHAYDIR9ydNAXnZSkAEEuBU7FCYF+WAgBxFACIfARvpCmwFAB4x1EAIPIR/AJpirOwFACIowBA5CM4OWmKe7IUAIijAEDkI+BT0tT1WAoAxFEAIPIRcAhpak6WAgBxFACIfAQsSZ76GUsBgDgKAEQ+AvimldOUpQBAHAUAIh8BTOToNGUpAJDbGGNsefHl9tTzah9WAC7h8fnLVywFAOIoAJApPjQf/WEFgFeWAgBxFADIFPloT9IUwHaWAgBxFADI8vloT9IUsBpLAYA4CgBEPpqUNAUcwVIAII4CAJGPLkqeAj5iKQAQRwGAyEfcnTQF52UpABBHAYDIR5yKNAX7shQAiKMAQOQjeCNNgaUAwDuOAgCRj+AXSFOchaUAQCwFODkrhHuyFACIowBA5CPgU9LU9VgKAMRRACDyEXAIaWpOlgIAcRQAiHwELEme+hlLAYA4CgBEPgL4ppXTlKUAQBwFACIfAUzk6DRlKQAQRwGA3MYYY8uLL7ennlf72g5wCY/PX75iKQAQRwGA/OjXR++/jt+bNAVwHEsBgDgKAGS6f16TpgCOYykAkOmWwp6sEID/z1IAII4CALlUPtqTNAWswFIAII4CAJGPTmDPNPXwIE8B/2MpABBHAYDIR/jlFBBLAYA4CgBEPmJX0hSci6UAQBwFACIfcVrSFNyfpQBAHAUAIh/BB6QprspSACCOAgCRj+CXSVPMzFIAII4CAJGPYCHSFP+UpQBALAVgEyvkGiwFAOIoABD5CDjcnmnq4UGe+g5LAYA4CgBEPgKW55dT21kKAMRRACDyEcA/sFqashQAiKMAQG5jjHH0HwHAHCwFAOIoABBHAYA4CgDEUQAgjgIAcRQAiKMAQBwFAPI3wbin372EhNIAAAAASUVORK5CYII=",280      "text/plain": [281       "<Figure size 640x480 with 1 Axes>"282      ]283     },284     "metadata": {},285     "output_type": "display_data"286    },287    {288     "data": {289      "image/png": "iVBORw0KGgoAAAANSUhEUgAAAYUAAAGFCAYAAAASI+9IAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjAsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvlHJYcgAAAAlwSFlzAAAPYQAAD2EBqD+naQAACMRJREFUeJzt3UGOI1UQBNAaxAoLcQEf0EeaA3IB1IJtseqQPQi5oX7l/1l6b8XCi/SopSCia6a+7fu+bwCwbdtPsw8AYB1CAYAQCgCEUAAghAIAIRQACKEAQAgFAOLnL3/yrz9OPAOgj8ftnv/+/ufvEy/5j3757e1HNAUAQigAEF+fjwD4h7ZT0r/QFAAIoQBACAWAQR63+8uc1JFQACCEAgDh6SOAwTo/kaQpABBCAYAwHwGcqNuUpCkAEJoCQJEOrUFTACCEAgBhPgKYYNUpSVMAIIQCAGE+AphspSlJUwAghAIAYT4CWMjsKUlTACCEAgBhPgJY1IwpSVMAIIQCAGE+AmigakrSFAAIoQBACAWAZh63+8ucNJJQACCEAgDh6SOAps54IklTACCEAgBhPgK4gFFTkqYAQGgKABdzpDVoCgCEUAAgzEcAF/YyJe0fbz+vKQAQQgGAMB8BHFD17uQqmgIAIRQACPMRwAFV706uoikAEEIBgBAKAIOc+e7kKkIBgBAKAISnjwAG6/xEkqYAQAgFAMJ8BHCiblOSpgBACAUAwnwEUKTDlKQpABBCAYAwHwFMsOqUpCkAEJoCwGQrtQZNAYAQCgCE+QhgIbOnJE0BgBAKAIT5CGBRM6YkTQGAEAoAhPkIoIGqKUlTACCEAgAhFACaedzuL3PSSEIBgBAKAISnjwCaOuOJJE0BgBAKAIT5COACRk1JmgIAIRQACPMRwMUcmZI0BQBCKAAQ5iOAC3uZkvaPt5/XFAAITQHggKp3J1fRFAAIoQBAmI8ADqh6d3IVTQGAEAoAhFAAGOTMdydXEQoAhFAAIDx9BDBY5yeSNAUAQigAEOYjgBN1m5I0BQBCKAAQ5iOAIh2mJE0BgBAKAIT5CGCCVackTQGAEAoAhPkIYLKVpiRNAYAQCgCE+QhgIbOnJE0BgBAKAIT5CGBRM6YkTQGA0BQAGqhqDZoCACEUAAihANDM43Z/mZNGEgoAhFAAIDx9BNDUGU8kaQoAhFAAIMxHABcwakrSFAAIoQBAmI8ALubIlKQpABBCAYAwHwFc2MuUtH+8/bymAEAIBQDCfARwQNW7k6toCgCEUAAgzEcAB5zxz1fPpCkAEEIBgBAKAIM8bveXOakjoQBA+EUzwGCdf/msKQAQQgGAMB8BnKjblKQpABBCAYAwHwEU6TAlaQoAhFAAIMxHABOsOiVpCgCEUAAgzEcAk600JWkKAIRQACDMRwALmT0laQoAhFAAIMxHAIuaMSVpCgCEUAAgzEcADVRNSZoCACEUAAihANDM43Z/mZNGEgoAhF80AzR1xi+fNQUAQigAEOYjgAsYNSVpCgCEUAAgzEcAF3NkStIUAAihAECYjwAu7GVK2j/efl5TACCEAgBhPgI4oOrdyVU0BQBCKAAQ5iOAA6renVxFUwAghAIAIRQABjnz3clVhAIAIRQACE8fAQzW+YkkTQGAEAoAhPkI4ETdpiRNAYDQFACKdGgNmgIAIRQACPMRwASrTkmaAgAhFAAI8xHAZCtNSZoCACEUAAjzEcBCZk9JmgIAIRQACPMRwKJmTEmaAgAhFAAI8xFAA1VTkqYAQAgFAEIoADTzuN1f5qSRhAIAIRQACE8fATR1xhNJmgIAIRQACPMRwAWMmpI0BQBCUwC4mCOtQVMAIIQCAGE+Ariwlylp/3j7eU0BgBAKAIT5COCAqncnV9EUAAihAECYjwAOqHp3chVNAYAQCgCEUAAY5Mx3J1cRCgCEUAAgPH0EMFjnJ5I0BQBCKAAQ5iOAE3WbkjQFAEIoABDmI4AiHaYkTQGAEAoAhPkIYIJVpyRNAYDQFAAmW6k1aAoAhFAAIMxHAAuZPSVpCgCEUAAgzEcAi5oxJWkKAIRQACDMRwANVE1JmgIAIRQACKEA0Mzjdn+Zk0YSCgCEUAAgPH0E0NQZTyRpCgCEUAAgzEcAFzBqStIUAAihAECYjwAu5siUpCkAEEIBgDAfAVzYy5S0f7z9vKYAQGgKAAdUvTu5iqYAQAgFAMJ8BHBA1buTq2gKAIRQACCEAsAgZ747uYpQACCEAgDh6SOAwTo/kaQpABBCAYAwHwGcqNuUpCkAEEIBgDAfARTpMCVpCgCEUAAgzEcAE6w6JWkKAIRQACDMRwCTrTQlaQoAhFAAIMxHAAuZPSVpCgCEUAAgzEcAi5oxJWkKAISmANBAVWvQFAAIoQBACAWAZh63+8ucNJJQACCEAgDh6SOAps54IklTACCEAgBhPgK4gFFTkqYAQAgFAMJ8BHAxR6YkTQGAEAoAhPkI4MJepqT94+3nNQUAQigAEOYjgAOq3p1cRVMAIIQCAPHl+ajqpdGjubte19vdXavr3T+6yvf4pCkAEEIBgPhfTx91rUvurtf1dnfX6nr3jz6/R+fvoCkAEIf/nkLXhHd3va7/F9X1z9zd83T+DpoCACEUAIih/8xF18rk7lpd7962vre7e55u30FTACCEAgBx2r+S6kmTWu6u52e8Vte7n3X4DpoCACEUAIjTX7LToS79m663u7tW17u3re/tXe9+tup30BQACKEAQJS+o3nVuvQVnjSp5e56fsbnWek7aAoAhFAAIErno2cr1aX/ouvd29b3dnfX6nr3tvW+/dPs76ApABBCAYCYNh89m12X/q+ud29b39vdXavr3dvW+/ZPM76DpgBACAUAYon56FnXytf17m3zl5aqubte15/xZ1V//poCACEUAIjl5qNnXeuqu+t1vd3dtbre/aMz5zBNAYBYuik865rw7q7X9ZeKXf/M3T3PGd9BUwAghAIA0WY+eta19rm7Vte7t63v7e6eZ9R30BQACKEAQLScj5550qSWu+v5Ga/V9e5nR76DpgBACAUA4tu+7/tXPvj49uvZtwzzXJeea1QHn7d3vXvbet3u7np+xuf5vn+8/YymAEAIBQDiy/MRANenKQAQQgGAEAoAhFAAIIQCACEUAAihAEAIBQBCKAAQfwNKUa9mB0hW1gAAAABJRU5ErkJggg==",290      "text/plain": [291       "<Figure size 640x480 with 1 Axes>"292      ]293     },294     "metadata": {},295     "output_type": "display_data"296    }297   ],298   "source": [299    "# Shuffling mask using torch.gather\n",300    "shuffle_index = torch.arange(num_tokens, dtype=torch.int64).reshape((-1, stride)).T.flip(0).flatten()\n",301    "mask1 = torch.gather(input=mask, dim=0, index=shuffle_index[:, None].expand(mask.shape))\n",302    "\n",303    "# plot_mask(mask1)\n",304    "plot_imshow(mask1)\n",305    "\n",306    "mask2 = torch.gather(input=mask1, dim=1, index=shuffle_index[None, :].expand(mask.shape))\n",307    "\n",308    "# plot_mask(mask2)\n",309    "plot_imshow(mask2)\n"310   ]311  },312  {313   "cell_type": "code",314   "execution_count": 3,315   "metadata": {},316   "outputs": [317    {318     "data": {319      "image/png": "iVBORw0KGgoAAAANSUhEUgAAAv8AAAL7CAYAAABqauo0AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjAsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvlHJYcgAAAAlwSFlzAAASdAAAEnQB3mYfeAAAErJJREFUeJzt3DFu68AZRtFRoEZb0F60eu/DnV25N2AbmNSBg8RKpCf/755TEiTxlRcDkIe9914AAMBf7x+PHgAAAPwZ4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEHK994PPt+d9eP50v//cYAADgOl8fLz++18k/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAxPFWL3p/fbrVq/6o0/ny6AkAAPBHOPkHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIi42d9+Jvw1Z+ofiQAA4Bac/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAEHHYe+9rHvh8e77XFq50Ol8ePQEAgAf7+nj58b1O/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIo63etGEP8+8vz59uzZ1NwAAXMvJPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEYe9977mgc+353ttIeJ0vjx6AgDAX+Pr4+XH9zr5BwCACPEPAAAR4h8AACLEPwAARIh/AACION7qRRP+4PL++vTt2tTda83eDgDAn+fkHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAARh733vuaBz7fne22BX+90vjx6AgDAv/j6ePnxvU7+AQAgQvwDAECE+AcAgAjxDwAAEcdbvWjCh5Dvr0/frk3dvdbc7VN3AwBM5+QfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABGHvfe+5oHPt+d7bQHu5HS+PHoCAHAnXx8vP77XyT8AAESIfwAAiBD/AAAQIf4BACBC/AMAQMTxVi+a8DeR99enb9em7l5r7vapu9eavR0AwMk/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAARh733vuaBz7fne20B+OZ0vjx6AgD8al8fLz++18k/AABEiH8AAIgQ/wAAECH+AQAgQvwDAEDE8VYvmvBHjvfXp2/Xpu5ea+72qbvXmrt96m4A4Lac/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAEHHYe+9rHvh8e77XFoC/xul8efQEACK+Pl5+fK+TfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCAiOOtXjThzxbvr0/frk3dvdbc7VN3rzV3+9Tda83eDgC/jZN/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAiDnvvfc0Dn2/P99oCwC9wOl8ePQGAK3x9vPz4Xif/AAAQIf4BACBC/AMAQIT4BwCACPEPAAARx1u9aMLfId5fn75dm7p7rbnbp+5ea+72qbvXmrt96m4A/m5O/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQcdh772se+Hx7vtcWAPifnc6XR08AeIivj5cf3+vkHwAAIsQ/AABEiH8AAIgQ/wAAEHG81YsmfGj1/vr07drU3WvN3T5191pzt0/dvdbc7VN3rzV7OwD/mZN/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAEQc9t77mgc+357vtQUAkk7ny6MnAIN9fbz8+F4n/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEcdbvWjCnwreX5++XZu6e62526fuXmvu9qm715q7feruteZun7ob4E9y8g8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQMRh772veeDz7fleWwCAQU7ny6MnAGutr4+XH9/r5B8AACLEPwAARIh/AACIEP8AABAh/gEAIOJ4qxdN+OL//fXp27Wpu9eau33q7rXmbp+6e62526fuXmvu9qm715q9HZjFyT8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABGHvfe+5oHPt+d7bQEA+CNO58ujJ8DNfH28/PheJ/8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHHW71owlfz769P365N3b3W3O1Td681d/vU3WvN3T5191pzt0/dvdbc7VN3Q5mTfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIg57733NA59vz/faAgDAf3A6Xx49gV/o6+Plx/c6+QcAgAjxDwAAEeIfAAAixD8AAESIfwAAiDje6kUTvj5/f336dm3q7rXmbp+6e62526fuXmvu9qm715q7feruteZun7p7rdnb4f/h5B8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEYe9977mgc+353ttAQDgL3U6Xx494a/19fHy43ud/AMAQIT4BwCACPEPAAAR4h8AACKOt3rRhI843l+fvl2bunutudun7l5r7vapu9eau33q7rXmbp+6e62526fuXmvu9qm7+T2c/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAg4rD33tc88Pn2fK8tAADwq5zOl0dP+K++Pl5+fK+TfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCAiOOtXjThS+j316dv16buXmvu9qm715q7feruteZun7p7rbnbp+5ea+72qbvXmrt96u61Zm//mzj5BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAg4rD33o8eAQAA3J+TfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARIh/AACIEP8AABAh/gEAIOKfb38K51byDvoAAAAASUVORK5CYII=",320      "text/plain": [321       "<Figure size 960x960 with 1 Axes>"322      ]323     },324     "metadata": {},325     "output_type": "display_data"326    }327   ],328   "source": [329    "import random\n",330    "import torch\n",331    "num_tokens = 128\n",332    "stride = 8\n",333    "mask = torch.zeros((num_tokens, num_tokens), dtype=torch.int32)\n",334    "\n",335    "iis = random.sample(range(0, num_tokens), k=10)\n",336    "\n",337    "# Adding diagonal elements with stride\n",338    "for i in range(0, num_tokens, stride):\n",339    "# for i in iis:\n",340    "    mask[i, i] = 1\n",341    "    mask[i:, i] = 1\n",342    "    mask[i, :i] = 1\n",343    "\n",344    "# js = random.sample(range(0, num_tokens), k=10)\n",345    "\n",346    "# for i in range(num_tokens):\n",347    "#     for j in range(0, i, stride):\n",348    "#     # for j in js:\n",349    "#         if i - j >= 0:\n",350    "#             mask[i, i - j] = 1\n",351    "\n",352    "plot_mask(mask)"353   ]354  },355  {356   "cell_type": "code",357   "execution_count": 4,358   "metadata": {},359   "outputs": [360    {361     "data": {362      "image/png": "iVBORw0KGgoAAAANSUhEUgAAAv8AAAL7CAYAAABqauo0AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjAsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvlHJYcgAAAAlwSFlzAAASdAAAEnQB3mYfeAAAFC9JREFUeJzt3UFqJUcaRtF040lvoffSq9c+PLNGnhe4DeoFKAIklBEvM+85w4erKGJ0+SE///bx8fFxAAAAj/evV/8DAACAPcQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiPj9u3/gf3/9Mfz93//574//MQAAwPf88/efX/5vXf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACK+vfYz8+v9bfi7FaCf87b7efP9vPl+3nw/b76fN9/Pm1+byz8AAESIfwAAiBD/AAAQIf4BACDitA9+Z0Yfffjg4xw+qNnPm+/nzffz5vt58/28+X7e/Bpc/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIpav/Yz42nstC0v7efP9vPl+3nw/b76fN9/Pm+/l8g8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAEPGStZ8ZK0DreNv9vPl+3nw/b76fN9/Pm+/nzddx+QcAgAjxDwAAEeIfAAAixD8AAESIfwAAiLjU2s+ML77X8bb7efP9vPl+3nw/b76fN9/Pm/+cyz8AAESIfwAAiBD/AAAQIf4BACDiFh/8zvjoYx1vu58338+b7+fN9/Pm+3nz/bz517n8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABE3HrtZ8YX3+t42/28+X7efD9vvp8338+b7+fNP3P5BwCACPEPAAAR4h8AACLEPwAARIh/AACIeOTaz8zoi+/y195n8jX9ft58P2++nzffz5vv5833K7+5yz8AAESIfwAAiBD/AAAQIf4BACBC/AMAQERq7Wek/LX3DhaW9vPm+3nz/bz5ft58P2++X+HNXf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACLyaz8zVoDW8bb7efP9vPl+3nw/b76fN9/vaW/u8g8AABHiHwAAIsQ/AABEiH8AAIjwwe83Pe2jjyvxtvt58/28+X7efD9vvp833++ub+7yDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQYe3nJHf94vsOvO1+3nw/b76fN9/Pm+/nzfe7+pu7/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARFj7WWz0xfdVvva+u6t/Tf9E3nw/b76fN9/Pm+/nzfe7ypu7/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARFj7eYGrfO39VBaW9vPm+3nz/bz5ft58P2++3+43d/kHAIAI8Q8AABHiHwAAIsQ/AABE/Pbx8fHxnT/wv7/+WPVvge1mH9TMPsrm57z5ft58P2++nzffz5vvN3vzf/7+88t/h8s/AABEiH8AAIgQ/wAAECH+AQAgQvwDAEDEaWs//tfPAACwn7UfAADgE/EPAAAR4h8AACLEPwAARIh/AACI+P2sv+jX+9vwdytAAABwDS7/AAAQIf4BACBC/AMAQIT4BwCACPEPAAARp639zIxWgCwAAQDAfi7/AAAQIf4BACBC/AMAQIT4BwCACPEPAAARy9d+RkYLQMdhBQgAAFZy+QcAgAjxDwAAEeIfAAAixD8AAES85IPfGR8CAwDAOi7/AAAQIf4BACBC/AMAQIT4BwCACPEPAAARl1r7mbECBAAAP+fyDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQcYu1nxkrQAAA8HUu/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEbde+5mxAgQAAJ+5/AMAQIT4BwCACPEPAAAR4h8AACIe+cHvzOhDYB8BAwBQ4fIPAAAR4h8AACLEPwAARIh/AACIEP8AABCRWvsZGS0AHYcVIAAAnsflHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgIr/2M2MFCACAp3H5BwCACPEPAAAR4h8AACLEPwAARIh/AACIsPbzTVaAAAC4K5d/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAIaz8nsQIEAMDVufwDAECE+AcAgAjxDwAAEeIfAAAifPC72OhDYB8BAwDwCi7/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR1n5eYLQAdBxWgAAAWMvlHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgwtrPhVgBAgBgJZd/AACIEP8AABAh/gEAIEL8AwBAhPgHAICI09Z+LNIAAMC1ufwDAECE+AcAgAjxDwAAEeIfAAAiTvvg99f72/B3HwIDAMA1uPwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESctvYzM1oBsgAEAAD7ufwDAECE+AcAgAjxDwAAEeIfAAAixD8AAEQsX/sZGS0AHYcVIAAAWMnlHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAg4iVrPzNWgAAAYB2XfwAAiBD/AAAQIf4BACBC/AMAQMSlPvid8SEwAAD8nMs/AABEiH8AAIgQ/wAAECH+AQAgQvwDAEDELdZ+ZqwAAQDA17n8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABE3HrtZ8YKEAAAfObyDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQ8ci1n5nRCpAFIAAAKlz+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAiUms/I6MFoOOwAgQAwPO4/AMAQIT4BwCACPEPAAAR4h8AACLyH/zO+BAYAICncfkHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIiw9vNNVoAAALgrl38AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAhrPyexAgQAwNW5/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARFj7WWy0AmQBCACAV3D5BwCACPEPAAAR4h8AACLEPwAARPjg9wVGHwEfhw+BAQBYy+UfAAAixD8AAESIfwAAiBD/AAAQIf4BACDC2s+FWAECAGAll38AAIgQ/wAAECH+AQAgQvwDAECE+AcAgIjT1n4s0gAAwLW5/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARJy29vPr/W34uxUgAAC4Bpd/AACIEP8AABAh/gEAIEL8AwBAhPgHAICI09Z+ZkYrQBaAAABgP5d/AACIEP8AABAh/gEAIEL8AwBAxPIPfkdGHwEfhw+BAQBgJZd/AACIEP8AABAh/gEAIEL8AwBAhPgHAICIl6z9zFgBAgCAdVz+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAiLrX2M2MFCAAAfs7lHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAg4hZrPzNWgAAA4Otc/gEAIEL8AwBAhPgHAIAI8Q8AABG3/uB3xofAAADwmcs/AABEiH8AAIgQ/wAAECH+AQAgQvwDAEDEI9d+ZkYrQBaAAACocPkHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIhIrf2MjBaAjsMKEAAAz+PyDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQkV/7mbECBADA07j8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEWPv5JitAAADclcs/AABEiH8AAIgQ/wAAECH+AQAgwge/J/EhMAAAV+fyDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQYe1nsdEKkAUgAABeweUfAAAixD8AAESIfwAAiBD/AAAQIf4BACDC2s8LjBaAjsMKEAAAa7n8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEWPu5ECtAAACs5PIPAAAR4h8AACLEPwAARIh/AACIOO2DXx+lAgDAtbn8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEnLb28+v9bfi7FSAAALgGl38AAIgQ/wAAECH+AQAgQvwDAECE+AcAgIjT1n5mRitAFoAAAGA/l38AAIgQ/wAAECH+AQAgQvwDAECE+AcAgIjlaz8jowWg47ACBAAAK7n8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEvGTtZ8YKEAAArOPyDwAAEeIfAAAixD8AAESIfwAAiLjUB78zPgQGAICfc/kHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIi4xdrPjBUgAAD4Opd/AACIEP8AABAh/gEAIEL8AwBAhPgHAICIW6/9zFgBAgCAz1z+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAiHrn2MzNaAbIABABAhcs/AABEiH8AAIgQ/wAAECH+AQAgIvXB78joI+Dj8CEwAADP4/IPAAAR4h8AACLEPwAARIh/AACIEP8AABCRX/uZsQIEAMDTuPwDAECE+AcAgAjxDwAAEeIfAAAixD8AAERY+/kmK0AAANyVyz8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIS1n5NYAQIA4Opc/gEAIEL8AwBAhPgHAIAI8Q8AABE++F1s9CGwj4ABAHgFl38AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAhrPy8wWgA6DitAAACs5fIPAAAR4h8AACLEPwAARIh/AACIEP8AABBh7edCrAABALCSyz8AAESIfwAAiBD/AAAQIf4BACBC/AMAQMRpaz8WaQAA4Npc/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIk5b+/n1/jb83QoQAABcg8s/AABEiH8AAIgQ/wAAECH+AQAg4rQPfmdGHwL7CBgAAPZz+QcAgAjxDwAAEeIfAAAixD8AAESIfwAAiFi+9jMyWgA6DitAAACwkss/AABEiH8AAIgQ/wAAECH+AQAgQvwDAEDES9Z+ZqwAAQDAOi7/AAAQIf4BACBC/AMAQIT4BwCACPEPAAARl1r7mbECBAAAP+fyDwAAEeIfAAAixD8AAESIfwAAiLjFB78zPgQGAICvc/kHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIi49drPjBUgAAD4zOUfAAAixD8AAESIfwAAiBD/AAAQIf4BACDikWs/M6MVIAtAAABUuPwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESk1n5GRgtAx2EFCACA53H5BwCACPEPAAAR4h8AACLEPwAARIh/AACIyK/9zFgBAgDgaVz+AQAgQvwDAECE+AcAgAjxDwAAET74/SYfAgMAcFcu/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEdZ+TmIFCACAq3P5BwCACPEPAAAR4h8AACLEPwAARIh/AACIsPaz2GgFyAIQAACv4PIPAAAR4h8AACLEPwAARIh/AACIEP8AABBh7ecFRgtAx2EFCACAtVz+AQAgQvwDAECE+AcAgAjxDwAAET74vRAfAgMAsJLLPwAARIh/AACIEP8AABAh/gEAIEL8AwBAxGlrPxZpAADg2lz+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAiTlv7+fX+NvzdChAAAFyDyz8AAESIfwAAiBD/AAAQIf4BACBC/AMAQMRpaz8zoxUgC0AAALCfyz8AAESIfwAAiBD/AAAQIf4BACBC/AMAQMTytZ+R0QLQcVgBAgCAlVz+AQAgQvwDAECE+AcAgAjxDwAAES/54HfGh8AAALCOyz8AAESIfwAAiBD/AAAQIf4BACBC/AMAQMSl1n5mrAABAMDPufwDAECE+AcAgAjxDwAAEeIfAAAixD8AAETcYu1nxgoQAAB8ncs/AABEiH8AAIgQ/wAAECH+AQAgQvwDAEDErdd+ZqwAAQDAZy7/AAAQIf4BACBC/AMAQIT4BwCAiEd+8Dsz+hDYR8AAAFS4/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARKTWfkZGC0DHYQUIAIDncfkHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIjIr/3MWAECAOBpXP4BACBC/AMAQIT4BwCACPEPAAAR4h8AACKs/XyTFSAAAO7K5R8AACLEPwAARIh/AACIEP8AABAh/gEAIMLaz0msAAEAcHUu/wAAECH+AQAgQvwDAECE+AcAgAgf/C42+hDYR8AAALyCyz8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIS1nxcYLQAdhxUgAADWcvkHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIiw9nMhVoAAAFjJ5R8AACLEPwAARIh/AACIEP8AABAh/gEAIOK0tR+LNAAAcG0u/wAAECH+AQAgQvwDAECE+AcAgIjTPvj99f42/N2HwAAAcA0u/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEaet/cyMVoAsAAEAwH4u/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEcvXfkZGC0DHYQUIAABWcvkHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIh4ydrPjBUgAABYx+UfAAAixD8AAESIfwAAiBD/AAAQcakPfmd8CAwAAD/n8g8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAEHGLtZ8ZK0AAAPB1Lv8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABG3XvuZsQIEAACfufwDAECE+AcAgAjxDwAAEeIfAAAixD8AAEQ8cu1nZrQCZAEIAIAKl38AAIgQ/wAAECH+AQAgQvwDAECE+AcAgIjU2s/IaAHoOKwAAQDwPC7/AAAQIf4BACBC/AMAQIT4BwCAiPwHvzM+BAYA4Glc/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIqz9fJMVIAAA7srlHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgwtrPSawAAQBwdS7/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR1n4WG60AWQACAOAVXP4BACBC/AMAQIT4BwCACPEPAAARPvh9gdFHwMfhQ2AAANZy+QcAgAjxDwAAEeIfAAAixD8AAESIfwAAiLD2cyFWgAAAWMnlHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAg4rePj4+PV/8jAACA9Vz+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCAiP8DO3bIPXU0nEQAAAAASUVORK5CYII=",363      "text/plain": [364       "<Figure size 960x960 with 1 Axes>"365      ]366     },367     "metadata": {},368     "output_type": "display_data"369    }370   ],371   "source": [372    "stride = 8\n",373    "\n",374    "shuffle_index = torch.arange(num_tokens, dtype=torch.int64).reshape((-1, stride)).T.flatten()\n",375    "# mask = mask[None, ...].expand(2, -1, -1)\n",376    "\n",377    "# Applying row-wise and column-wise shuffling\n",378    "mask1 = torch.gather(mask, dim=0, index=shuffle_index[:, None].expand(mask.shape))\n",379    "mask2 = torch.gather(mask1, dim=1, index=shuffle_index[None, :].expand(mask.shape))\n",380    "\n",381    "plot_mask(mask2)\n",382    "# plot_imshow(mask2[0,0])\n",383    "\n",384    "# def shuffle_mask(mask, stride):\n",385    "#     shuffle_index = torch.arange(num_tokens, dtype=torch.int64).reshape((-1, stride)).T.flatten()\n",386    "#     mask1 = torch.gather(mask, dim=0, index=shuffle_index[:, None].expand(mask.shape))\n",387    "#     mask2 = torch.gather(mask1, dim=1, index=shuffle_index[None, :].expand(mask.shape))\n",388    "#     return mask2"389   ]390  },391  {392   "cell_type": "code",393   "execution_count": null,394   "metadata": {},395   "outputs": [],396   "source": [397    "shuffle_index"398   ]399  },400  {401   "cell_type": "code",402   "execution_count": null,403   "metadata": {},404   "outputs": [],405   "source": [406    "import numpy as np\n",407    "from glob import glob\n",408    "\n",409    "flops_files = glob('plots/extra_analysis/longvila_flops_counter_*.txt')\n",410    "flops_files.sort(key=lambda x: int(x.split('_')[-1].split('.')[0]))\n",411    "\n",412    "for flops_file in flops_files:\n",413    "    ctx_len = flops_file.split('_')[-1].split('.')[0]\n",414    "    with open(flops_file, 'r') as f:\n",415    "        flops_data = [float(line.strip()) for line in f.readlines()]\n",416    "    print(f\"ctx_len: {ctx_len}, avg ratio: {np.mean(flops_data)}\")"417   ]418  },419  {420   "cell_type": "code",421   "execution_count": null,422   "metadata": {},423   "outputs": [],424   "source": [425    "import numpy as np\n",426    "from glob import glob\n",427    "\n",428    "flops_files = glob('plots/extra_analysis/longvila_flops_counter_*.txt')\n",429    "flops_files.sort(key=lambda x: int(x.split('_')[-1].split('.')[0]))\n",430    "\n",431    "for flops_file in flops_files:\n",432    "    ctx_len = flops_file.split('_')[-1].split('.')[0]\n",433    "    with open(flops_file, 'r') as f:\n",434    "        flops_data = [float(line.strip()) for line in f.readlines()]\n",435    "    print(f\"ctx_len: {ctx_len}, avg ratio: {np.mean(flops_data)}\")"436   ]437  },438  {439   "cell_type": "code",440   "execution_count": 5,441   "metadata": {},442   "outputs": [],443   "source": [444    "import torch\n",445    "\n",446    "def create_diagonal_pattern(size=128, stride=8):\n",447    "    indices = torch.arange(size)\n",448    "    x, y = torch.meshgrid(indices, indices, indexing='ij')\n",449    "    pattern = (((x - y) % stride == 0) & (x >= y)).float()\n",450    "    \n",451    "    return pattern\n",452    "\n",453    "# Create the pattern\n",454    "pattern = create_diagonal_pattern(size=128, stride=8)\n"455   ]456  },457  {458   "cell_type": "code",459   "execution_count": 6,460   "metadata": {},461   "outputs": [462    {463     "data": {464      "image/png": "iVBORw0KGgoAAAANSUhEUgAAAv8AAAL7CAYAAABqauo0AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjAsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvlHJYcgAAAAlwSFlzAAASdAAAEnQB3mYfeAAAHLBJREFUeJzt3b2OZWlahNEYlJSEj0V53AdXz32UB1b5SBRSY4wQP9/eTJ/pzJMnMtYyj7paj7YVKul760+//fbbbwEAAL68v/nsAAAA4DmMfwAAGGH8AwDACOMfAABGGP8AADDC+AcAgBHGPwAAjDD+AQBghPEPAAAjjH8AABhh/AMAwIi3R//Ar58/Ln//u3/4pz8cAwAAPOY//v1ffvd/62/+AQBghPEPAAAjjH8AABhh/AMAwAjjHwAARvzpt99+++2RP/D27fvl7//2r/98+bsrQAAA8HFc+wEAAA7GPwAAjDD+AQBghPEPAAAj3u3B752rh8AeAQMAwPvw4BcAADgY/wAAMML4BwCAEcY/AACMMP4BAGDEh1/7uXJ1AShxBQgAAB7l2g8AAHAw/gEAYITxDwAAI4x/AAAYYfwDAMCIT7n2c8cVIAAAeIxrPwAAwMH4BwCAEcY/AACMMP4BAGCE8Q8AACNe6trPHVeAAADgmms/AADAwfgHAIARxj8AAIww/gEAYETFg987HgIDALDOg18AAOBg/AMAwAjjHwAARhj/AAAwwvgHAIAR1dd+7rgCBADACtd+AACAg/EPAAAjjH8AABhh/AMAwAjjHwAARjx87efXzx+Xvzdc0rm6AtTQDQAAd1z7AQAADsY/AACMMP4BAGCE8Q8AACOMfwAAGPHwtZ+3b98vf7+6pJO8/jWd1m4AAEhc+wEAAC4Y/wAAMML4BwCAEcY/AACMMP4BAGDEu137uXN1Tafhko4rQAAANHDtBwAAOBj/AAAwwvgHAIARxj8AAIz48Ae/V5of0za3AwDw9XjwCwAAHIx/AAAYYfwDAMAI4x8AAEYY/wAAMOJTrv3cab6k09wOAEAv134AAICD8Q8AACOMfwAAGGH8AwDACOMfAABGvNS1nzvNl3Su2hu6AQDo4NoPAABwMP4BAGCE8Q8AACOMfwAAGGH8AwDAiIprP3darwC1dgMA8Hpc+wEAAA7GPwAAjDD+AQBghPEPAAAjqh/83ml9UNvaDQDA5/HgFwAAOBj/AAAwwvgHAIARxj8AAIww/gEAYMTD135+/fxx+XvDRZqrazqt3UlHOwAAH8u1HwAA4GD8AwDACOMfAABGGP8AADDC+AcAgBEPX/t5+/b98vfWizSt3Ul3OwAA78O1HwAA4GD8AwDACOMfAABGGP8AADDC+AcAgBHvdu3nztVFmoZrNM2XdFq/OQAAj3PtBwAAOBj/AAAwwvgHAIARxj8AAIww/gEAYMSHX/u58tUu6SSv397aDQDA/8+1HwAA4GD8AwDACOMfAABGGP8AADDiUx783ml+lNra3toNAMCfefALAAAcjH8AABhh/AMAwAjjHwAARhj/AAAw4qWu/dxpvkhz1d7anXS0AwAsce0HAAA4GP8AADDC+AcAgBHGPwAAjDD+AQBgRMW1nzutF2lau5PudgCAr8i1HwAA4GD8AwDACOMfAABGGP8AADDC+AcAgBHV137utF6kae1OutsBAJq59gMAAByMfwAAGGH8AwDACOMfAABGPPzg99fPH5e/NzzsvHqU2tqd9LY3dAMAtPDgFwAAOBj/AAAwwvgHAIARxj8AAIww/gEAYMTD137evn2//L31Ik1rd9Lb3toNAPCKXPsBAAAOxj8AAIww/gEAYITxDwAAI4x/AAAY8W7Xfu5cXXZpuOrSfJHGNwcA2OHaDwAAcDD+AQBghPEPAAAjjH8AABhh/AMAwIgPv/ZzpfmqS2t7a3fS3Q4A8NFc+wEAAA7GPwAAjDD+AQBghPEPAAAjjH8AABjxKdd+7jRfdWltb+1OutsBAN6Laz8AAMDB+AcAgBHGPwAAjDD+AQBgxEs9+L3T/LDzqr21O+ltb+gGAPhrePALAAAcjH8AABhh/AMAwAjjHwAARhj/AAAwouLaz53WizSt3Ulve2s3AMBf4toPAABwMP4BAGCE8Q8AACOMfwAAGGH8AwDAiOprP3daL7u0die97a3dAAD/xbUfAADgYPwDAMAI4x8AAEYY/wAAMML4BwCAEQ9f+/n188fl7w3XUa4uu7R2J73trd1JRzsAsMW1HwAA4GD8AwDACOMfAABGGP8AADDi4Qe/b9++X/7e+kCytTvpbW/tTrrbAYCvyYNfAADgYPwDAMAI4x8AAEYY/wAAMML4BwCAEe927efO1XWUhssozVddfPPna/3mAEA/134AAICD8Q8AACOMfwAAGGH8AwDACOMfAABGfPi1nytf7apL8vrtrd1Jb3trNwDQxbUfAADgYPwDAMAI4x8AAEYY/wAAMML4BwCAEZ9y7edO83WU1vbW7qS3vbUbAHhNrv0AAAAH4x8AAEYY/wAAMML4BwCAES/14PdO8wPJq/bW7qS3vbU76WgHAD6PB78AAMDB+AcAgBHGPwAAjDD+AQBghPEPAAAjKq793Gm9jtLanfS2t3Yn3e0AwMdz7QcAADgY/wAAMML4BwCAEcY/AACMMP4BAGBE9bWfO63XUVq7k9721u6kux0AeD+u/QAAAAfjHwAARhj/AAAwwvgHAIARxj8AAIx4+NrPr58/Ln9vuDBydR2ltTvpbW/tTnrbG7oBgL+Oaz8AAMDB+AcAgBHGPwAAjDD+AQBghPEPAAAjHr728/bt++XvrddRWruT3vbW7qS3vbUbAPjLXPsBAAAOxj8AAIww/gEAYITxDwAAI97twe+dq4eGDY8Mmx9I+ubP55sDAJ/Fg18AAOBg/AMAwAjjHwAARhj/AAAwwvgHAIARH37t50rzhZHW9tbupLe9tTvpbgeANa79AAAAB+MfAABGGP8AADDC+AcAgBHGPwAAjPiUaz93mi+MtLa3die97a3dSXc7AHxVrv0AAAAH4x8AAEYY/wAAMML4BwCAEcY/AACMeKlrP3eaL4xctbd2J73trd1Jb3tDNwB8Ba79AAAAB+MfAABGGP8AADDC+AcAgBEVD37vtD6QbO1Oettbu5Pe9tZuAGjjwS8AAHAw/gEAYITxDwAAI4x/AAAYYfwDAMCI6ms/d1qvjLR2J73trd1Jb3trNwC8Ktd+AACAg/EPAAAjjH8AABhh/AMAwAjjHwAARjx87efXzx+Xvzdc6ri6MtLanfS2t3Ynve2t3UlHOwB8Jtd+AACAg/EPAAAjjH8AABhh/AMAwAjjHwAARjx87eft2/fL31svdbR2J73trd1Jb3trd9LdDgDP4NoPAABwMP4BAGCE8Q8AACOMfwAAGGH8AwDAiHe79nPn6lJHw5WO5gsjvvnz+ebP1/rNAeC9ufYDAAAcjH8AABhh/AMAwAjjHwAARnz4g98rX+2RYfL67a3dSW97a3fS297aDQB/hAe/AADAwfgHAIARxj8AAIww/gEAYITxDwAAIz7l2s+d5ksdre2t3Ulve2t30tve2g0Av4drPwAAwMH4BwCAEcY/AACMMP4BAGCE8Q8AACNe6trPneZLHVftrd1Jb3trd9Lb3tqddLQDwH9x7QcAADgY/wAAMML4BwCAEcY/AACMMP4BAGBExbWfO62XOlq7k9721u6kt721O+luB2CPaz8AAMDB+AcAgBHGPwAAjDD+AQBgRPWD3zutj/Vau5Pe9tbupLe9tTvpbgfg6/LgFwAAOBj/AAAwwvgHAIARxj8AAIww/gEAYMTD135+/fxx+XvDtYurSx2t3Ulve2t30tve2p30tjd0A/A1uPYDAAAcjH8AABhh/AMAwAjjHwAARhj/AAAw4uFrP2/fvl/+3nqpo7U76W1v7U5621u7k9721m4A+rj2AwAAHIx/AAAYYfwDAMAI4x8AAEYY/wAAMOLdrv3cubp40XDtovlSh2/+fL758/nmAPBnrv0AAAAH4x8AAEYY/wAAMML4BwCAEcY/AACM+PBrP1ear120trd2J73trd1Jb3trd9LdDsDncu0HAAA4GP8AADDC+AcAgBHGPwAAjPiUB793mh+8tba3die97a3dSW97a3fS3Q7Ac3jwCwAAHIx/AAAYYfwDAMAI4x8AAEYY/wAAMOKlrv3cab52cdXe2p30trd2J73trd1Jb3tDNwDvz7UfAADgYPwDAMAI4x8AAEYY/wAAMML4BwCAERXXfu60Xupo7U5621u7k9721u6kt721G4A/xrUfAADgYPwDAMAI4x8AAEYY/wAAMML4BwCAEdXXfu60Xrxo7U5621u7k9721u6kt721G4Dfx7UfAADgYPwDAMAI4x8AAEYY/wAAMOLhB7+/fv64/L3h4djVo7fW7qS3vbU76W1v7U5621u7k452AP6bB78AAMDB+AcAgBHGPwAAjDD+AQBghPEPAAAjHr728/bt++XvrVcjWruT3vbW7qS3vbU76W1v7U662wEWufYDAAAcjH8AABhh/AMAwAjjHwAARhj/AAAw4t2u/dy5uhrRcDGi+dqFb/58vvnz+ebP1/rNAb46134AAICD8Q8AACOMfwAAGGH8AwDACOMfAABGfPi1nytf7dpF8vrtrd1Jb3trd9Lb3tqd9La3dgN8Ja79AAAAB+MfAABGGP8AADDC+AcAgBHGPwAAjPiUaz93mq9GtLa3die97a3dSW97a3fS297aDdDItR8AAOBg/AMAwAjjHwAARhj/AAAw4qUe/N5pfjh21d7anfS2t3Ynve2t3Ulve2t30tEO8Ko8+AUAAA7GPwAAjDD+AQBghPEPAAAjjH8AABhRce3nTuvViNbupLe9tTvpbW/tTnrbW7uT7naAz+baDwAAcDD+AQBghPEPAAAjjH8AABhh/AMAwIjqaz93Wq9GtHYnve2t3Ulve2t30tve2p10twM8i2s/AADAwfgHAIARxj8AAIww/gEAYITxDwAAIx6+9vPr54/L3xsuL1xdjWjtTnrbW7uT3vbW7qS3vbU76W1v6Ab4CK79AAAAB+MfAABGGP8AADDC+AcAgBEPP/h9+/b98vfWh2Ot3Ulve2t30tve2p30trd2J73trd0Af5QHvwAAwMH4BwCAEcY/AACMMP4BAGCE8Q8AACPe7drPndZ/gr35aoRv/ny++fP55s/nmwO8Jtd+AACAg/EPAAAjjH8AABhh/AMAwAjjHwAARnz4tZ8rzZcXWttbu5Pe9tbupLe9tTvpbW/tTrrbAf4n134AAICD8Q8AACOMfwAAGGH8AwDACOMfAABGfMq1nzvNlxda21u7k9721u6kt721O+ltb+1OutuBTa79AAAAB+MfAABGGP8AADDC+AcAgBEv9eD3TvPjq6v21u6kt721O+ltb+1Oettbu5Pe9oZu4Ovz4BcAADgY/wAAMML4BwCAEcY/AACMMP4BAGBExbWfO61XI1q7k9721u6kt721O+ltb+1Oettbu4GvxbUfAADgYPwDAMAI4x8AAEYY/wAAMML4BwCAEdXXfu60Xl9o7U5621u7k9721u6kt721O+ltb+0GOrn2AwAAHIx/AAAYYfwDAMAI4x8AAEYY/wAAMOLhaz+/fv64/L3hgsHV9YXW7qS3vbU76W1v7U5621u7k9721u6kox14Xa79AAAAB+MfAABGGP8AADDC+AcAgBHGPwAAjHj42s/bt++Xv7deMGjtTnrbW7uT3vbW7qS3vbU76W1v7U6624HP59oPAABwMP4BAGCE8Q8AACOMfwAAGPFuD37v+CfYn883fz7f/Pl88+fzzZ+v9ZsDz+XBLwAAcDD+AQBghPEPAAAjjH8AABhh/AMAwIgPv/Zz5atdXkhev721O+ltb+1Oettbu5Pe9tbupLe9tRv4OK79AAAAB+MfAABGGP8AADDC+AcAgBHGPwAAjPiUaz93mi8YtLa3die97a3dSW97a3fS297anfS2t3YDf5xrPwAAwMH4BwCAEcY/AACMMP4BAGCE8Q8AACNe6trPneYLBlftrd1Jb3trd9Lb3tqd9La3die97a3dSUc78Pu49gMAAByMfwAAGGH8AwDACOMfAABGVDz4vdP6iKm1O+ltb+1Oettbu5Pe9tbupLe9tTvpbgf+Nw9+AQCAg/EPAAAjjH8AABhh/AMAwAjjHwAARlRf+7nTesGgtTvpbW/tTnrbW7uT3vbW7qS3vbU76W6HVa79AAAAB+MfAABGGP8AADDC+AcAgBHGPwAAjHj42s+vnz8uf2+4AnB1waC1O+ltb+1Oettbu5Pe9tbupLe9tTvpbW/ohgWu/QAAAAfjHwAARhj/AAAwwvgHAIARxj8AAIx4+NrP27fvl7+3XjBo7U5621u7k9721u6kt721O+ltb+1Oettbu+Grce0HAAA4GP8AADDC+AcAgBHGPwAAjDD+AQBgxLtd+7lzdQmg4QpA8wUD3/z5fPPn882fzzd/Pt8c+D1c+wEAAA7GPwAAjDD+AQBghPEPAAAjPvzB75Xmh0Ct7a3dSW97a3fS297anfS2t3Ynve2t3Ul3O7wyD34BAICD8Q8AACOMfwAAGGH8AwDACOMfAABGfMq1nzvNVwBa21u7k9721u6kt721O+ltb+1Oettbu5PudngFrv0AAAAH4x8AAEYY/wAAMML4BwCAEcY/AACMeKlrP3earwBctbd2J73trd1Jb3trd9Lb3tqd9La3die97Q3d8Gyu/QAAAAfjHwAARhj/AAAwwvgHAIARxj8AAIyouPZzp/WCQWt30tve2p30trd2J73trd1Jb3trd9Lb3toNH8m1HwAA4GD8AwDACOMfAABGGP8AADCi+sHvndbHQK3dSW97a3fS297anfS2t3Ynve2t3Ulve2s3vAcPfgEAgIPxDwAAI4x/AAAYYfwDAMAI4x8AAEY8fO3n188fl783vKa/ugTQ2p30trd2J73trd1Jb3trd9Lb3tqd9La3dicd7fB7ufYDAAAcjH8AABhh/AMAwAjjHwAARhj/AAAw4uFrP2/fvl/+3vqavrU76W1v7U5621u7k9721u6kt721O+ltb+1Outvh/3LtBwAAOBj/AAAwwvgHAIARxj8AAIww/gEAYMS7Xfu5c/WavuElffMVAN/8+Xzz5/PNn883fz7f/PlavznbXPsBAAAOxj8AAIww/gEAYITxDwAAI4x/AAAY8eHXfq58tSsAyeu3t3Ynve2t3Ulve2t30tve2p30trd2J73trd3scO0HAAA4GP8AADDC+AcAgBHGPwAAjPiUB793mh/UtLa3die97a3dSW97a3fS297anfS2t3Ynve2t3Xw9HvwCAAAH4x8AAEYY/wAAMML4BwCAEcY/AACMeKlrP3eaX9Nftbd2J73trd1Jb3trd9Lb3tqd9La3die97a3dSUc7nVz7AQAADsY/AACMMP4BAGCE8Q8AACOMfwAAGFFx7edO62v61u6kt721O+ltb+1Oettbu5Pe9tbupLe9tTvpbue1ufYDAAAcjH8AABhh/AMAwAjjHwAARhj/AAAwovraz53W1/St3Ulve2t30tve2p30trd2J73trd1Jb3trd9Ldzmtw7QcAADgY/wAAMML4BwCAEcY/AACMePjB76+fPy5/b3iUcvWgprU76W1v7U5621u7k9721u6kt721O+ltb+1Oetsbunk+D34BAICD8Q8AACOMfwAAGGH8AwDACOMfAABGPHzt5+3b98vfW1/Tt3Ynve2t3Ulve2t30tve2p30trd2J73trd1Jb3trNx/LtR8AAOBg/AMAwAjjHwAARhj/AAAwwvgHAIAR73bt587Vq/SGF+nNr+l98+fzzZ/PN38+3/z5fPPn881p5NoPAABwMP4BAGCE8Q8AACOMfwAAGGH8AwDAiA+/9nOl+UV6a3trd9Lb3tqd9La3die97a3dSW97a3fS297anXS38/u59gMAAByMfwAAGGH8AwDACOMfAABGGP8AADDiU6793Gl+kd7a3tqd9La3die97a3dSW97a3fS297anfS2t3Yn3e2cXPsBAAAOxj8AAIww/gEAYITxDwAAI17qwe+d5kcpV+2t3Ulve2t30tve2p30trd2J73trd1Jb3trd9Lb3tC9zoNfAADgYPwDAMAI4x8AAEYY/wAAMML4BwCAERXXfu60vqZv7U5621u7k9721u6kt721O+ltb+1Oettbu5Pe9tbuJa79AAAAB+MfAABGGP8AADDC+AcAgBHGPwAAjKi+9nOn9VV6a3fS297anfS2t3Ynve2t3Ulve2t30tve2p30trd2f0Wu/QAAAAfjHwAARhj/AAAwwvgHAIARxj8AAIx4+NrPr58/Ln9veNl99Sq9tTvpbW/tTnrbW7uT3vbW7qS3vbU76W1v7U5621u7k472Vq79AAAAB+MfAABGGP8AADDC+AcAgBEPP/h9+/b98vfWxx2t3Ulve2t30tve2p30trd2J73trd1Jb3trd9Lb3tqddLe/Og9+AQCAg/EPAAAjjH8AABhh/AMAwAjjHwAARrzbtZ87/mnq5/PNn883fz7f/Pl88+fzzZ/PN3++1m/+Slz7AQAADsY/AACMMP4BAGCE8Q8AACOMfwAAGPHh136ufLUX6cnrt7d2J73trd1Jb3trd9Lb3tqd9La3die97a3dSW97a/dnce0HAAA4GP8AADDC+AcAgBHGPwAAjDD+AQBgxKdc+7nT/LK7tb21O+ltb+1Oettbu5Pe9tbupLe9tTvpbW/tTnrbW7s/mms/AADAwfgHAIARxj8AAIww/gEAYMRLPfi90/y446q9tTvpbW/tTnrbW7uT3vbW7qS3vbU76W1v7U5621u7k4729+DBLwAAcDD+AQBghPEPAAAjjH8AABhh/AMAwIiKaz93Wl92t3Ynve2t3Ulve2t30tve2p30trd2J73trd1Jb3trd9Ld/gjXfgAAgIPxDwAAI4x/AAAYYfwDAMAI4x8AAEZUX/u50/qyu7U76W1v7U5621u7k9721u6kt721O+ltb+1Oettbu5Pu9iuu/QAAAAfjHwAARhj/AAAwwvgHAIARxj8AAIx4+NrPr58/Ln9veB199bK7tTvpbW/tTnrbW7uT3vbW7qS3vbU76W1v7U5621u7k972hm7XfgAAgIPxDwAAI4x/AAAYYfwDAMAI4x8AAEY8fO3n7dv3y99bX3a3die97a3dSW97a3fS297anfS2t3Ynve2t3Ulve2t30tve0O3aDwAAcDD+AQBghPEPAAAjjH8AABjxbg9+77T+M8kNjzvu+ObP55s/n2/+fL758/nmz+ebP59v/sd58AsAAByMfwAAGGH8AwDACOMfAABGGP8AADDiw6/9XHml19GPam1v7U5621u7k9721u6kt721O+ltb+1Oettbu5Pe9tbu5HPaXfsBAAAOxj8AAIww/gEAYITxDwAAI4x/AAAY8SnXfu542f18rd1Jb3trd9Lb3tqd9La3die97a3dSW97a3fS297anXxsu2s/AADAwfgHAIARxj8AAIww/gEAYITxDwAAI17q2s+dr/ayu7U76W1v7U5621u7k9721u6kt721O+ltb+1Oettbu5Pe9ke7XfsBAAAOxj8AAIww/gEAYITxDwAAIyoe/N5pfdzR2p30trd2J73trd1Jb3trd9Lb3tqd9La3die97a3dSW/7o90e/AIAAAfjHwAARhj/AAAwwvgHAIARxj8AAIyovvZzZ+Vl9ytpbW/tTnrbW7uT3vbW7qS3vbU76W1v7U5621u7k972u+6//ft//N3/D3/zDwAAI4x/AAAYYfwDAMAI4x8AAEYY/wAAMOLhaz8AAEAnf/MPAAAjjH8AABhh/AMAwAjjHwAARhj/AAAwwvgHAIARxj8AAIww/gEAYITxDwAAI4x/AAAYYfwDAMCI/wSnR0Vk9dPRagAAAABJRU5ErkJggg==",465      "text/plain": [466       "<Figure size 960x960 with 1 Axes>"467      ]468     },469     "metadata": {},470     "output_type": "display_data"471    }472   ],473   "source": [474    "q_len = 128\n",475    "stride = 8\n",476    "\n",477    "indices = torch.arange(q_len)\n",478    "x, y = torch.meshgrid(indices, indices, indexing='ij')\n",479    "mask = (((x - y) % stride == 0) & (x >= y)).float()\n",480    "# mask[::stride,:] = 1\n",481    "# mask[:,::stride] = 1\n",482    "\n",483    "plot_mask(mask)\n",484    "\n"485   ]486  },487  {488   "cell_type": "code",489   "execution_count": 7,490   "metadata": {},491   "outputs": [492    {493     "data": {494      "image/png": "iVBORw0KGgoAAAANSUhEUgAAAv8AAAL7CAYAAABqauo0AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjAsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvlHJYcgAAAAlwSFlzAAASdAAAEnQB3mYfeAAAEpRJREFUeJzt3L2OI9cVRtErYWxAuSM783vo6fUezqRIuQFZQDvQr33IGbKbRVbVXiskSIDhxgXO99Xb29vbAgAATu/rV/8BAADgOcQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiPh07w/+8+O/Ln7+zd+//fCfAQAA7vPzT9/f/F0v/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEXev/Vzz7x++u/i5FSAAANgHL/8AABAh/gEAIEL8AwBAhPgHAICIhx38XnPpENgRMAAAPJ+XfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCAiM3Xfi65tAC0lhUgAADYkpd/AACIEP8AABAh/gEAIEL8AwBAhPgHAICIl6z9XGMFCAAAtuPlHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgYldrP9dYAQIAgI/z8g8AABHiHwAAIsQ/AABEiH8AAIg4xMHvNQ6BAQDgdl7+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAiDr32c40VIAAAmLz8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEnHLt55pLK0AWgAAAqPDyDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQkVr7ueTSAtBaVoAAADgfL/8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABFfvb29vd3zg09//cfFz6+t5pyNFSAAAPbk55++v/m7Xv4BACBC/AMAQIT4BwCACPEPAAARDzv4vcYhMAAAbMfBLwAAMIh/AACIEP8AABAh/gEAIEL8AwBAxOZrP5dUFoDWsgIEAMC2rP0AAACD+AcAgAjxDwAAEeIfAAAixD8AAES8ZO3nmsoKkAUgAAAexdoPAAAwiH8AAIgQ/wAAECH+AQAgQvwDAEDErtZ+rrECBAAAl1n7AQAABvEPAAAR4h8AACLEPwAARBzi4Pcah8AAANQ5+AUAAAbxDwAAEeIfAAAixD8AAESIfwAAiDj02s81VoAAAKiw9gMAAAziHwAAIsQ/AABEiH8AAIgQ/wAAEHHKtZ9rrAABAHA21n4AAIBB/AMAQIT4BwCACPEPAAAR4h8AACJSaz+XWAACAODIrP0AAACD+AcAgAjxDwAAEeIfAAAixD8AAETk136usQIEAMARWPsBAAAG8Q8AABHiHwAAIsQ/AABEOPi9k0NgAAD2xMEvAAAwiH8AAIgQ/wAAECH+AQAgQvwDAECEtZ8HsQIEAMArWPsBAAAG8Q8AABHiHwAAIsQ/AABEiH8AAIiw9rMxK0AAAGzJ2g8AADCIfwAAiBD/AAAQIf4BACBC/AMAQIS1nxeoLACtZQUIAGBr1n4AAIBB/AMAQIT4BwCACPEPAAARDn53pHII7AgYAOBxHPwCAACD+AcAgAjxDwAAEeIfAAAixD8AAERY+zkAK0AAAFxj7QcAABjEPwAARIh/AACIEP8AABAh/gEAIMLaz4FZAQIAwNoPAAAwiH8AAIgQ/wAAECH+AQAgQvwDAECEtZ8TsgIEANBh7QcAABjEPwAARIh/AACIEP8AABAh/gEAIMLaT0RlAWgtK0AAQIu1HwAAYBD/AAAQIf4BACBC/AMAQISD37jKIbAjYADgrBz8AgAAg/gHAIAI8Q8AABHiHwAAIsQ/AABEWPvhIitAAADHYO0HAAAYxD8AAESIfwAAiBD/AAAQIf4BACDC2g93sQIEALAv1n4AAIBB/AMAQIT4BwCACPEPAAAR4h8AACKs/fAQVoAAAF7D2g8AADCIfwAAiBD/AAAQIf4BACDCwS+bcggMALAtB78AAMAg/gEAIEL8AwBAhPgHAIAI8Q8AABHWfng6C0AAAI9j7QcAABjEPwAARIh/AACIEP8AABAh/gEAIMLaD7thBQgA4H7WfgAAgEH8AwBAhPgHAIAI8Q8AABHiHwAAIqz9sHtWgAAArrP2AwAADOIfAAAixD8AAESIfwAAiHDwy2E5BAYAcPALAABcIP4BACBC/AMAQIT4BwCACPEPAAAR1n44HStAAECJtR8AAGAQ/wAAECH+AQAgQvwDAECE+AcAgAhrPyRUFoDWsgIEADXWfgAAgEH8AwBAhPgHAIAI8Q8AABHiHwAAIqz9kFZZAbIABADnZe0HAAAYxD8AAESIfwAAiBD/AAAQIf4BACDC2g9cYAUIADgKaz8AAMAg/gEAIEL8AwBAhPgHAIAIB79wB4fAAMDeOPgFAAAG8Q8AABHiHwAAIsQ/AABEiH8AAIiw9gMPYAUIAHgVaz8AAMAg/gEAIEL8AwBAhPgHAIAI8Q8AABHWfmAjlQWgtawAAcArWfsBAAAG8Q8AABHiHwAAIsQ/AABEiH8AAIiw9gNPVlkBsgAEAM9h7QcAABjEPwAARIh/AACIEP8AABDh4Bd2wiEwAPAeDn4BAIBB/AMAQIT4BwCACPEPAAAR4h8AACKs/cDOWQECAD7H2g8AADCIfwAAiBD/AAAQIf4BACBC/AMAQIS1HzgoK0AAwFrWfgAAgAvEPwAARIh/AACIEP8AABAh/gEAIMLaD5yMFSAAaLH2AwAADOIfAAAixD8AAESIfwAAiBD/AAAQYe0HAiwAAcB5WfsBAAAG8Q8AABHiHwAAIsQ/AABEOPiFMIfAAHB8Dn4BAIBB/AMAQIT4BwCACPEPAAAR4h8AACKs/QCDFSAAOA5rPwAAwCD+AQAgQvwDAECE+AcAgAjxDwAAEdZ+gJtZAQKA/bH2AwAADOIfAAAixD8AAESIfwAAiBD/AAAQYe0H+DArQADwOtZ+AACAQfwDAECE+AcAgAjxDwAAEQ5+gU1UjoDXcggMwGs5+AUAAAbxDwAAEeIfAAAixD8AAESIfwAAiLD2AzxVZQXIAhAAz2LtBwAAGMQ/AABEiH8AAIgQ/wAAECH+AQAgwtoPsAtWgADgfaz9AAAAg/gHAIAI8Q8AABHiHwAAIsQ/AABEWPsBds0KEAB8nrUfAABgEP8AABAh/gEAIEL8AwBAhPgHAIAIaz/AIVkBAoBfWPsBAAAG8Q8AABHiHwAAIsQ/AABEOPgFTqNyBLyWQ2AA/uDgFwAAGMQ/AABEiH8AAIgQ/wAAECH+AQAgwtoPcHqVFSALQABN1n4AAIBB/AMAQIT4BwCACPEPAAAR4h8AACKs/QBZVoAAOANrPwAAwCD+AQAgQvwDAECE+AcAgAjxDwAAEdZ+AP6PFSAAjsTaDwAAMIh/AACIEP8AABAh/gEAIMLBL8CNHAIDsEcOfgEAgEH8AwBAhPgHAIAI8Q8AABHiHwAAIqz9AHyQFSAAXsnaDwAAMIh/AACIEP8AABAh/gEAIEL8AwBAhLUfgA1YAALgWaz9AAAAg/gHAIAI8Q8AABHiHwAAIsQ/AABEWPsBeCIrQAA8mrUfAABgEP8AABAh/gEAIEL8AwBAhPgHAIAIaz8AO2AFCID3svYDAAAM4h8AACLEPwAARIh/AACIcPALsGMOgQH4Ege/AADAIP4BACBC/AMAQIT4BwCACPEPAAAR1n4ADsgKEAC/sfYDAAAM4h8AACLEPwAARIh/AACIEP8AABBh7QfgJCoLQGtZAQL4M2s/AADAIP4BACBC/AMAQIT4BwCACPEPAAAR1n4ATq6yAmQBCKiy9gMAAAziHwAAIsQ/AABEiH8AAIhw8AsQ5RAY4Bwc/AIAAIP4BwCACPEPAAAR4h8AACLEPwAARFj7AeB/WAECOBZrPwAAwCD+AQAgQvwDAECE+AcAgAjxDwAAEdZ+ALiJFSCAfbL2AwAADOIfAAAixD8AAESIfwAAiBD/AAAQYe0HgHerLACtZQUI2C9rPwAAwCD+AQAgQvwDAECE+AcAgAgHvwA8XOUQ2BEwsAcOfgEAgEH8AwBAhPgHAIAI8Q8AABHiHwAAIqz9APA0VoAAHs/aDwAAMIh/AACIEP8AABAh/gEAIEL8AwBAhLUfAF7OChDA+1n7AQAABvEPAAAR4h8AACLEPwAARIh/AACIsPYDwG5ZAQL4Mms/AADAIP4BACBC/AMAQIT4BwCACPEPAAAR1n4AOBwrQAB/sPYDAAAM4h8AACLEPwAARIh/AACIcPALwCk4AgaqHPwCAACD+AcAgAjxDwAAEeIfAAAixD8AAERY+wHg1KwAAWdn7QcAABjEPwAARIh/AACIEP8AABAh/gEAIMLaDwBJVoCAs7D2AwAADOIfAAAixD8AAESIfwAAiBD/AAAQYe0HAP7EChBwNNZ+AACAQfwDAECE+AcAgAjxDwAAEQ5+AeAGDoGBvXLwCwAADOIfAAAixD8AAESIfwAAiBD/AAAQYe0HAN6psgC0lhUg2DNrPwAAwCD+AQAgQvwDAECE+AcAgAjxDwAAEdZ+AODBKitAFoBgH6z9AAAAg/gHAIAI8Q8AABHiHwAAIsQ/AABEWPsBgCexAgRswdoPAAAwiH8AAIgQ/wAAECH+AQAgQvwDAECEtR8AeDErQMBHWPsBAAAG8Q8AABHiHwAAIsQ/AABEOPgFgJ1yCAzcwsEvAAAwiH8AAIgQ/wAAECH+AQAgQvwDAECEtR8AOJDKAtBaVoDgVtZ+AACAQfwDAECE+AcAgAjxDwAAEeIfAAAirP0AwAlUVoAsAMFk7QcAABjEPwAARIh/AACIEP8AABAh/gEAIMLaDwCcmBUgOD9rPwAAwCD+AQAgQvwDAECE+AcAgAgHvwAQ5BAYzsPBLwAAMIh/AACIEP8AABAh/gEAIEL8AwBAhLUfAOB3VoDgeKz9AAAAg/gHAIAI8Q8AABHiHwAAIsQ/AABEWPsBAL7IChDsl7UfAABgEP8AABAh/gEAIEL8AwBAhPgHAIAIaz8AwLtYAIJ9sPYDAAAM4h8AACLEPwAARIh/AACIEP8AABBh7QcAeCgrQPBc1n4AAIBB/AMAQIT4BwCACPEPAAARDn4BgKdwCAzbcPALAAAM4h8AACLEPwAARIh/AACIEP8AABBh7QcAeCkrQPAx1n4AAIBB/AMAQIT4BwCACPEPAAAR4h8AACKs/QAAu2QFCG5j7QcAABjEPwAARIh/AACIEP8AABAh/gEAIMLaDwBwGJUFoLWsAHE7az8AAMAg/gEAIEL8AwBAhPgHAIAIB78AwOFVDoEdAXOJg18AAGAQ/wAAECH+AQAgQvwDAECE+AcAgAhrPwDAaVkBosDaDwAAMIh/AACIEP8AABAh/gEAIEL8AwBAhLUfACDHChBnYu0HAAAYxD8AAESIfwAAiBD/AAAQIf4BACDC2g8AwK+sAHFE1n4AAIBB/AMAQIT4BwCACPEPAAAR4h8AACKs/QAAfEZlAWgtK0BHZe0HAAAYxD8AAESIfwAAiBD/AAAQ4eAXAOAdKofAjoD3z8EvAAAwiH8AAIgQ/wAAECH+AQAgQvwDAECEtR8AgAeyAsSzWfsBAAAG8Q8AABHiHwAAIsQ/AABEiH8AAIiw9gMA8ARWgNiKtR8AAGAQ/wAAECH+AQAgQvwDAECE+AcAgAhrPwAAL2QFiI+y9gMAAAziHwAAIsQ/AABEiH8AAIhw8AsAsEMOgbmVg18AAGAQ/wAAECH+AQAgQvwDAECE+AcAgAhrPwAAB2EBiEus/QAAAIP4BwCACPEPAAAR4h8AACLEPwAARFj7AQA4OCtAbdZ+AACAQfwDAECE+AcAgAjxDwAAEeIfAAAirP0AAJyUFaAGaz8AAMAg/gEAIEL8AwBAhPgHAIAIB78AADEOgc/FwS8AADCIfwAAiBD/AAAQIf4BACBC/AMAQIS1HwAA1lpWgI7K2g8AADCIfwAAiBD/AAAQIf4BACBC/AMAQIS1HwAArqosAK113BUgaz8AAMAg/gEAIEL8AwBAhPgHAIAI8Q8AABHWfgAAuFtlBegIC0DWfgAAgEH8AwBAhPgHAIAI8Q8AABHiHwAAIqz9AADwMFaAns/aDwAAMIh/AACIEP8AABAh/gEAIMLBLwAAm3MIvB0HvwAAwCD+AQAgQvwDAECE+AcAgAjxDwAAEdZ+AAB4GStAH2ftBwAAGMQ/AABEiH8AAIgQ/wAAECH+AQAgwtoPAAC7UlkAWusxK0DWfgAAgEH8AwBAhPgHAIAI8Q8AABHiHwAAIqz9AABwCJUVoHsXgKz9AAAAg/gHAIAI8Q8AABHiHwAAIhz8AgBwaPVDYAe/AADAIP4BACBC/AMAQIT4BwCACPEPAAAR1n4AADilygrQX/72z5u/6+UfAAAixD8AAESIfwAAiBD/AAAQIf4BACDi7rUfAADgmLz8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEf8F9U3F71eIN00AAAAASUVORK5CYII=",495      "text/plain": [496       "<Figure size 960x960 with 1 Axes>"497      ]498     },499     "metadata": {},500     "output_type": "display_data"501    }502   ],503   "source": [504    "q_len = 128\n",505    "stride = 8\n",506    "\n",507    "indices = torch.arange(q_len)\n",508    "x, y = torch.meshgrid(indices, indices, indexing='ij')\n",509    "mask = (((x - y) < 10) & (x >= y)).float()\n",510    "# mask[::stride,:] = 1\n",511    "# mask[:,::stride] = 1\n",512    "\n",513    "plot_mask(mask)\n",514    "\n"515   ]516  },517  {518   "cell_type": "code",519   "execution_count": 8,520   "metadata": {},521   "outputs": [522    {523     "data": {524      "image/png": "iVBORw0KGgoAAAANSUhEUgAAAv8AAAL7CAYAAABqauo0AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjAsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvlHJYcgAAAAlwSFlzAAASdAAAEnQB3mYfeAAAG/5JREFUeJzt3cGqHUl6hdEscy3wvEfWzO/hp+/30Kx6pLmhVHD9AJEBir4Z+2SevdbwoGqab5BsBPHrj8/Pz88DAAB4e//x6v8DAABAhvEPAAAljH8AAChh/AMAQAnjHwAAShj/AABQwvgHAIASxj8AAJQw/gEAoITxDwAAJYx/AAAo8bH6H/z6+eP09//67//98v8ZAABgzd9//fnbf9bf/AMAQAnjHwAAShj/AABQwvgHAIASxj8AAJRYvvYz83//+ufp764AfZ22eZrnaZ6neZ7meZrnaX5v/uYfAABKGP8AAFDC+AcAgBLGPwAAlPjj8/Pzc+U/+Pj2/fT32eOOMx58XMODmjzN8zTP0zxP8zzN91nZhMeh+RX+/uvP3/6z/uYfAABKGP8AAFDC+AcAgBLGPwAAlDD+AQCgxGXXfmZcAco7a67tXprnaZ6neZ7meZrvYxPu49oPAAAwMP4BAKCE8Q8AACWMfwAAKGH8AwBAie3Xfs6svPY+Di++rzBrru0+mudpnqd5nuZ5mu/lCtDXufYDAAAMjH8AAChh/AMAQAnjHwAAShj/AABQ4iXXfmZcAcpzwSBP8zzN8zTP0zxP831swjWu/QAAAAPjHwAAShj/AABQwvgHAIASt3rwO+PRR55HTHma52mep3me5nma72MTnvPgFwAAGBj/AABQwvgHAIASxj8AAJQw/gEAoMQjrv3MePGd54JBnuZ5mudpnqd5nub7tG9C134AAICB8Q8AACWMfwAAKGH8AwBACeMfAABKPPraz8zKi+93e+39Ki4Y5Gmep3me5nma52m+T8sVINd+AACAgfEPAAAljH8AAChh/AMAQAnjHwAASrzltZ8ZV4Dyzppru5fmeZrnaZ6neZ7m+7zbJnTtBwAAGBj/AABQwvgHAIASxj8AAJQw/gEAoETVtZ8zK6+9j+MZL77vbtZc2300z9M8T/M8zfM03+upV4Bc+wEAAAbGPwAAlDD+AQCghPEPAAAl6h/8zngInOcRU57meZrnaZ6neZ7m+zxhE3rwCwAADIx/AAAoYfwDAEAJ4x8AAEoY/wAAUMK1n0VPePH9blwwyNM8T/M8zfM0z9N8nzttQtd+AACAgfEPAAAljH8AAChh/AMAQAnjHwAASrj2c5GVF99e2F/DBYM8zfM0z9M8T/M8zfd5xRUg134AAICB8Q8AACWMfwAAKGH8AwBACeMfAABKuPazmStAeWfNtd1L8zzN8zTP0zxP8312bkLXfgAAgIHxDwAAJYx/AAAoYfwDAEAJD35f4BX/7HM7/4x5nuZ5mudpnqd5nuZ7XfEQ2INfAABgYPwDAEAJ4x8AAEoY/wAAUML4BwCAEsvXfn79/HH6uxff+3hlDwDAjGs/AADAwPgHAIASxj8AAJQw/gEAoITxDwAAJT6u+h9ykSZP8320zdM8T/M8zfM0z9P83vzNPwAAlDD+AQCghPEPAAAljH8AAChh/AMAQIk/Pj8/P1f+g49v309/n73sPuO19zU0z3PBIE/zPM3zNM/TPE/zff7+68/f/rP+5h8AAEoY/wAAUML4BwCAEsY/AACUMP4BAKDEZdd+ZlykydM876y5tntpnqd5nuZ5mudp/nWu/QAAAAPjHwAAShj/AABQwvgHAIAS2x/8nll5kHocHn1cQfM8/4x5nuZ5mudpnqd5nuZrPPgFAAAGxj8AAJQw/gEAoITxDwAAJYx/AAAo8ZJrPzMu0uRpnueCQZ7meZrnaZ6neZ7m51z7AQAABsY/AACUMP4BAKCE8Q8AACWMfwAAKHGraz8zLtLkaZ7ngkGe5nma52mep3lee3PXfgAAgIHxDwAAJYx/AAAoYfwDAEAJ4x8AAEo84trPjIs0eZrntV8weAXN8zTP0zxP87yW5q79AAAAA+MfAABKGP8AAFDC+AcAgBKPfvA7s/Io9d0efLyK5nktj5juRPM8zfM0z9M8792ae/ALAAAMjH8AAChh/AMAQAnjHwAAShj/AABQ4i2v/cy4SJOned5Zc2330jxP8zzN8zTPe2pz134AAICB8Q8AACWMfwAAKGH8AwBACeMfAABKVF37ObNyjeY4nvHi++40z5s113YfzfM0z9M8T/O8JzR37QcAABgY/wAAUML4BwCAEsY/AACUMP4BAKBE/bWfGRdp8jTPe8IFg3ejeZ7meZrnaZ53p+au/QAAAAPjHwAAShj/AABQwvgHAIASxj8AAJRw7WeRizR5mufd6YJBC83zNM/TPE/zvFc0d+0HAAAYGP8AAFDC+AcAgBLGPwAAlPDg9yIrj1I9srmG5nkejuVpnqd5nuZ5muftbO7BLwAAMDD+AQCghPEPAAAljH8AAChh/AMAQAnXfjZzkSZP87yz5trupXme5nma52med0Vz134AAICB8Q8AACWMfwAAKGH8AwBACeMfAABKuPbzAivXaI7DK/sraJ43a67tPprnaZ6neZ7meavNXfsBAAAGxj8AAJQw/gEAoITxDwAAJYx/AAAosXzt59fPH6e/e/G9j1f2eZoDAE/h2g8AADAw/gEAoITxDwAAJYx/AAAosfzg9+Pb99PfPZDcZ9Z2RvOv0zzPNyRP8zzN8zTP0zzPg18AAGBg/AMAQAnjHwAAShj/AABQwvgHAIASl137mTl78e219zVWLtJofg3N81yNyNM8T/M8zfM038e1HwAAYGD8AwBACeMfAABKGP8AAFDC+AcAgBLbr/2c8dp7Lxdp8jTPc0ksT/M8zfM0z9P861z7AQAABsY/AACUMP4BAKCE8Q8AACWMfwAAKPGSaz8zrgDts3KN5jg0v4Lmeb4heZrnaZ6neZ7ma1z7AQAABsY/AACUMP4BAKCE8Q8AACVu9eB3xqOPfTxKzdM8zzckT/M8zfM0z9P8nAe/AADAwPgHAIASxj8AAJQw/gEAoITxDwAAJR5x7WfGi+99XKTJ0zzPNyRP8zzN8zTPa2/u2g8AADAw/gEAoITxDwAAJYx/AAAoYfwDAECJR1/7mWl/8b2TizR5muf5huRpnqd5nuZ5Lc1d+wEAAAbGPwAAlDD+AQCghPEPAAAljH8AACjxltd+Zs5efL/ba+9XWblIo/k1NM9ruRpxJ5rnaZ6ned67NXftBwAAGBj/AABQwvgHAIASxj8AAJQw/gEAoETVtZ8z7/ba+25cpMnTPM8lsTzN8zTP0zzvqc1d+wEAAAbGPwAAlDD+AQCghPEPAAAl6h/8zngIvM/Kg9Tj0PwKmuf5huRpnqd5nuZ5T2juwS8AADAw/gEAoITxDwAAJYx/AAAoYfwDAEAJ134WPeHF91O5SJOneZ5vSJ7meZrnaZ53p+au/QAAAAPjHwAAShj/AABQwvgHAIASxj8AAJRw7ecid3rx/W5cpMnTPM83JE/zPM3zNM97RXPXfgAAgIHxDwAAJYx/AAAoYfwDAEAJ4x8AAEq49rPZ2YtvL+yvsXKRRvNraJ7nUkee5nma52met7O5az8AAMDA+AcAgBLGPwAAlDD+AQCghAe/L+CRzV4epeZpnueYQJ7meZrnaZ53RXMPfgEAgIHxDwAAJYx/AAAoYfwDAEAJ4x8AAEq49nMjrgDts3KN5jg0v4Lmeb4heZrnaZ6ned5qc9d+AACAgfEPAAAljH8AAChh/AMAQAnjHwAASixf+/n188fp71587+OVfZ7meZoDwL/HtR8AAGBg/AMAQAnjHwAAShj/AABQwvgHAIASy9d+Pr59P/3dpY48zfeZtZ3R/Os0z/MNydM8T/M8zfNc+wEAAAbGPwAAlDD+AQCghPEPAAAljH8AAChx2bWfmbMX315776X5PisXaTS/huZ5LnXkaZ6neZ7m+7j2AwAADIx/AAAoYfwDAEAJ4x8AAEpsf/B7xoOPPM338ig1T/M8xwTyNM/TPE/zr/PgFwAAGBj/AABQwvgHAIASxj8AAJQw/gEAoMRLrv3MuEiTp/k+K9dojkPzK2ie5xuSp3me5nmar3HtBwAAGBj/AABQwvgHAIASxj8AAJQw/gEAoMStrv3MePGdp/k+LtLkaZ7nG5KneZ7meZqfc+0HAAAYGP8AAFDC+AcAgBLGPwAAlDD+AQCgxCOu/cx48Z2n+T4u0uRpnucbkqd5nuZ57c1d+wEAAAbGPwAAlDD+AQCghPEPAAAlHv3gd6b90ccraL6PR6l5muf5huRpnqd5XktzD34BAICB8Q8AACWMfwAAKGH8AwBACeMfAABKvOW1n5mzF9/v9tr7bjTfZ+UijebX0Dyv5VLHnWiep3neuzV37QcAABgY/wAAUML4BwCAEsY/AACUMP4BAKBE1bWfM+/22vsJNN/LRZo8zfNcEsvTPE/zvKc2d+0HAAAYGP8AAFDC+AcAgBLGPwAAlDD+AQCgRP21nxkXafI032flGs1xaH4FzfN8Q/I0z9M87wnNXfsBAAAGxj8AAJQw/gEAoITxDwAAJYx/AAAo4drPoie8+H43mu/jIk2e5nm+IXma52med6fmrv0AAAAD4x8AAEoY/wAAUML4BwCAEh78XuROjz5aaL6PR6l5muf5huRpnqd53iuae/ALAAAMjH8AAChh/AMAQAnjHwAAShj/AABQwrWfzc5efHthv5fm+6xcpNH8GprnuY6Sp3me5nk7m7v2AwAADIx/AAAoYfwDAEAJ4x8AAEoY/wAAUMK1nxfwwj5P871cpMnTPM8lsTzN8zTPu6K5az8AAMDA+AcAgBLGPwAAlDD+AQCghPEPAAAlXPu5ERdp8jTfZ+UazXFofgXN83xD8jTP0zxvtblrPwAAwMD4BwCAEsY/AACUMP4BAKDE8oPfXz9/nP7u0cc+HtrkaZ6neZ7mAO/Bg18AAGBg/AMAQAnjHwAAShj/AABQwvgHAIASy9d+Pr59P/3d1Yg8zfM032fWdkbzr9M8zzckT/M8zfNc+wEAAAbGPwAAlDD+AQCghPEPAAAljH8AAChx2bWfmbMX315776V5nub7rFyk0fwamue5jpKneZ7m+7j2AwAADIx/AAAoYfwDAEAJ4x8AAEoY/wAAUGL7tZ8zXnvnaZ6n+V4u0uRpnueSWJ7meZp/nWs/AADAwPgHAIASxj8AAJQw/gEAoITxDwAAJV5y7WfGdZQ8zfM032flGs1xaH4FzfN8Q/I0z9N8jWs/AADAwPgHAIASxj8AAJQw/gEAoMStHvzOePSRp3me5vt4lJqneZ5vSJ7meZqf8+AXAAAYGP8AAFDC+AcAgBLGPwAAlDD+AQCgxCOu/cx48Z2neZ7m+7hIk6d5nm9InuZ57c1d+wEAAAbGPwAAlDD+AQCghPEPAAAljH8AACjx6Gs/M+0vvl9B8zzN93GRJk/zPN+QPM3zWpq79gMAAAyMfwAAKGH8AwBACeMfAABKGP8AAFDiLa/9zJy9+H631953o3me5vusXKTR/Bqa57VcR7kTzfPerblrPwAAwMD4BwCAEsY/AACUMP4BAKBE1YPfM+/24OMJNM/TfC+PUvM0z3NMIE/zvKc29+AXAAAYGP8AAFDC+AcAgBLGPwAAlDD+AQCgRP21nxnXUfI0z9N8n5VrNMeh+RU0z/MNydM87wnNXfsBAAAGxj8AAJQw/gEAoITxDwAAJYx/AAAo4drPoie8+H43mudpvo+LNHma5/mG5Gmed6fmrv0AAAAD4x8AAEoY/wAAUML4BwCAEsY/AACUcO3nInd68d1C8zzN93GRJk/zPN+QPM3zXtHctR8AAGBg/AMAQAnjHwAAShj/AABQwoPfzc4efXhks5fmeZrvs/IoVfNraJ7nUWqe5nk7m3vwCwAADIx/AAAoYfwDAEAJ4x8AAEoY/wAAUMK1nxfwwj5P8zzN93KRJk/zPJfE8jTPu6K5az8AAMDA+AcAgBLGPwAAlDD+AQCghPEPAAAlXPu5EddR8jTP03yflWs0x6H5FTTP8w3J0zxvtblrPwAAwMD4BwCAEsY/AACUMP4BAKCE8Q8AACWWr/38+vnj9Hcvvvfxyj5P8zzN8zTP0xzYwbUfAABgYPwDAEAJ4x8AAEoY/wAAUML4BwCAEsvXfj6+fT/93QWDPM3zNM/TfJ9Z2xnNv07zPN+QPM3zXPsBAAAGxj8AAJQw/gEAoITxDwAAJS578Dtz9ujDg4+9NM/TPE/zfVYepWp+Dc3zPErN03wfD34BAICB8Q8AACWMfwAAKGH8AwBACeMfAABKbL/2c8Zr7zzN8zTP03wvF2nyNM9zSSxP869z7QcAABgY/wAAUML4BwCAEsY/AACUMP4BAKDES679zLjUkad5nuZ5mu+zco3mODS/guZ5viF5mq9x7QcAABgY/wAAUML4BwCAEsY/AACUMP4BAKDEra79zHjxnad5nuZ5mu/jIk2e5nm+IXman3PtBwAAGBj/AABQwvgHAIASxj8AAJR4xIPfGY8+8jTP0zxP8308Ss3TPM83JK+9uQe/AADAwPgHAIASxj8AAJQw/gEAoITxDwAAJR597Wem/cX3K2iep3me5vu4SJOneZ5vSF5Lc9d+AACAgfEPAAAljH8AAChh/AMAQAnjHwAASrzltZ+Zsxff7/ba+240z9M8T/N9Vi7SaH4NzfNaLtLcybs1d+0HAAAYGP8AAFDC+AcAgBLGPwAAlDD+AQCgRNW1nzPv9tr7CTTP0zxP871cpMnTPM8lsbynNnftBwAAGBj/AABQwvgHAIASxj8AAJQw/gEAoET9tZ8ZlzryNM/TPE/zfVau0RyH5lfQPM83JO8JzV37AQAABsY/AACUMP4BAKCE8Q8AACU8+F30hEcf70bzPM3zNN/Ho9Q8zfN8Q/Lu1NyDXwAAYGD8AwBACeMfAABKGP8AAFDC+AcAgBKu/VzkTi++W2iep3me5vu4SJOneZ5vSN4rmrv2AwAADIx/AAAoYfwDAEAJ4x8AAEoY/wAAUMK1n83OXnx7Yb+X5nma52m+z8pFGs2voXmeK0B5O5u79gMAAAyMfwAAKGH8AwBACeMfAABKGP8AAFDCtZ8X8MI+T/M8zfM038tFmjzN81wSy7uiuWs/AADAwPgHAIASxj8AAJQw/gEAoIQHvzfisV6e5nma52m+z8qD1OPQ/Aqa5/mG5K029+AXAAAYGP8AAFDC+AcAgBLGPwAAlDD+AQCgxPK1n18/f5z+7sX3Pl7Z52mep3me5nma52lOA9d+AACAgfEPAAAljH8AAChh/AMAQAnjHwAASixf+/n49v30d6/p8zTP0zxP8zzN95m1ndH86zTP8w3Jc+0HAAAYGP8AAFDC+AcAgBLGPwAAlDD+AQCgxGXXfmbOXnx77b2X5nma52mep/k+KxdpNL+G5nmuAO3j2g8AADAw/gEAoITxDwAAJYx/AAAoYfwDAECJ7dd+znjtnad5nuZ5mudpvpeLNHma57kk9nWu/QAAAAPjHwAAShj/AABQwvgHAIASL3nwO+PhWJ7meZrnaZ6n+T4rD1KPQ/MraJ7nG7LGg18AAGBg/AMAQAnjHwAAShj/AABQwvgHAIASt7r2M+PFd57meZrnaZ6n+T4u0uRpnucbcs61HwAAYGD8AwBACeMfAABKGP8AAFDC+AcAgBKPuPYz48V3nuZ5mudpnqf5Pi7S5Gme1/4Nce0HAAAYGP8AAFDC+AcAgBLGPwAAlDD+AQCgxKOv/cy0v/h+Bc3zNM/TPE/zfVykydM8r+Ub4toPAAAwMP4BAKCE8Q8AACWMfwAAKPGWD35nzh59vNuDj7vRPE/zPM3zNN9n5VGq5tfQPO/dHgJ78AsAAAyMfwAAKGH8AwBACeMfAABKGP8AAFCi6trPmXd77f0Emudpnqd5nuZ7uUiTp3neUy+JufYDAAAMjH8AAChh/AMAQAnjHwAAShj/AABQov7az4yrEXma52mep3me5vusXKM5Ds2voHneE74hrv0AAAAD4x8AAEoY/wAAUML4BwCAEsY/AACUcO1n0RNefL8bzfM0z9M8T/N9XKTJ0zzvTt8Q134AAICB8Q8AACWMfwAAKGH8AwBACeMfAABKuPZzkTu9+G6heZ7meZrnab6PizR5mue94hvi2g8AADAw/gEAoITxDwAAJYx/AAAo4cHvZmePPjym2UvzPM3zNM/TfJ+VR6maX0PzvJ0PgT34BQAABsY/AACUMP4BAKCE8Q8AACWMfwAAKOHazwv4p+PzNM/TPE/zPM33cpEmT/O8Ky6JufYDAAAMjH8AAChh/AMAQAnjHwAAShj/AABQwrWfG3E1Ik/zPM3zNM/TfJ+VazTHofkVNM9b/Ya49gMAAAyMfwAAKGH8AwBACeMfAABKGP8AAFBi+drPr58/Tn/3snsfVyPyNM/TPE/zPM3zNM/TPM+1HwAAYGD8AwBACeMfAABKGP8AAFBi+cHvx7fvp7973JGneZ7meZrnaZ6n+T6ztjOaf53meR78AgAAA+MfAABKGP8AAFDC+AcAgBLGPwAAlLjs2s/M2Ytvr7r30jxP8zzN8zTP03yflYs0ml9D831c+wEAAAbGPwAAlDD+AQCghPEPAAAljH8AACix/drPmdlrby+799E8T/M8zfM0z9N8Lxdp8jT/Otd+AACAgfEPAAAljH8AAChh/AMAQAnjHwAASrzk2s+MCwZ5mudpnqd5nuZ5mu+zco3mODS/guZrXPsBAAAGxj8AAJQw/gEAoITxDwAAJW714HfGI6Y8zfM0z9M8T/M8zffxKDVP83Me/AIAAAPjHwAAShj/AABQwvgHAIASxj8AAJR4xLWfGRcM8jTP0zxP8zzN8zTfx0WavPbmrv0AAAAD4x8AAEoY/wAAUML4BwCAEsY/AACUePS1nxkXDPI0z9M8T/M8zfM036f9Is0rtDR37QcAABgY/wAAUML4BwCAEsY/AACUMP4BAKDEW177mTl78f3UV91PoXme5nma52mep/k+KxdpNL/GuzV37QcAABgY/wAAUML4BwCAEsY/AACUMP4BAKBE1bWfM7PX3k942f1Umudpnqd5nuZ5mu/1bhdpnuCpzV37AQAABsY/AACUMP4BAKCE8Q8AACXqH/zOeMSUp3me5nma52mep/k+Kw9Sj0PzKzyhuQe/AADAwPgHAIASxj8AAJQw/gEAoITxDwAAJVz7WeSCQZ7meZrnaZ6neZ7m+zzhIs27uVNz134AAICB8Q8AACWMfwAAKGH8AwBACeMfAABKuPZzERcM8jTP0zxP8zzN8zTf504XaVq8orlrPwAAwMD4BwCAEsY/AACUMP4BAKCE8Q8AACVc+9ns7MW3l/R7aZ6neZ7meZrnab7PykUaza+xs7lrPwAAwMD4BwCAEsY/AACUMP4BAKCEB78v4J8xz9M8T/M8zfM0z9N8Lw+B865o7sEvAAAwMP4BAKCE8Q8AACWMfwAAKGH8AwBACdd+bsQFgzzN8zTP0zxP8zzN91m5RnMcml9htfl//uN/fvvP+pt/AAAoYfwDAEAJ4x8AAEoY/wAAUML4BwCAEsvXfgAAgGfyN/8AAFDC+AcAgBLGPwAAlDD+AQCghPEPAAAljH8AAChh/AMAQAnjHwAAShj/AABQwvgHAIASxj8AAJT4fz7RI2Pbt4YPAAAAAElFTkSuQmCC",525      "text/plain": [526       "<Figure size 960x960 with 1 Axes>"527      ]528     },529     "metadata": {},530     "output_type": "display_data"531    }532   ],533   "source": [534    "stride = 8\n",535    "\n",536    "shuffle_index = torch.arange(num_tokens, dtype=torch.int64).reshape((-1, stride)).T.flatten()\n",537    "# mask = mask[None, ...].expand(2, -1, -1)\n",538    "\n",539    "# Applying row-wise and column-wise shuffling\n",540    "mask1 = torch.gather(mask, dim=0, index=shuffle_index[:, None].expand(mask.shape))\n",541    "mask2 = torch.gather(mask1, dim=1, index=shuffle_index[None, :].expand(mask.shape))\n",542    "\n",543    "plot_mask(mask2)"544   ]545  },546  {547   "cell_type": "code",548   "execution_count": 9,549   "metadata": {},550   "outputs": [551    {552     "ename": "FileNotFoundError",553     "evalue": "[Errno 2] No such file or directory: 'mminference_best_patterns.json'",554     "output_type": "error",555     "traceback": [556      "\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",557      "\u001b[0;31mFileNotFoundError\u001b[0m                         Traceback (most recent call last)",558      "Cell \u001b[0;32mIn[9], line 6\u001b[0m\n\u001b[1;32m      2\u001b[0m \u001b[38;5;28;01mfrom\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;21;01mcollections\u001b[39;00m\u001b[38;5;250m \u001b[39m\u001b[38;5;28;01mimport\u001b[39;00m Counter\n\u001b[1;32m      4\u001b[0m seq_len \u001b[38;5;241m=\u001b[39m \u001b[38;5;241m20020\u001b[39m\n\u001b[0;32m----> 6\u001b[0m \u001b[38;5;28;01mwith\u001b[39;00m \u001b[38;5;28;43mopen\u001b[39;49m\u001b[43m(\u001b[49m\u001b[38;5;124;43m\"\u001b[39;49m\u001b[38;5;124;43mmminference_best_patterns.json\u001b[39;49m\u001b[38;5;124;43m\"\u001b[39;49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;124;43m\"\u001b[39;49m\u001b[38;5;124;43mr\u001b[39;49m\u001b[38;5;124;43m\"\u001b[39;49m\u001b[43m)\u001b[49m \u001b[38;5;28;01mas\u001b[39;00m f:\n\u001b[1;32m      7\u001b[0m     data \u001b[38;5;241m=\u001b[39m json\u001b[38;5;241m.\u001b[39mload(f)\n\u001b[1;32m      9\u001b[0m \u001b[38;5;28;01mwith\u001b[39;00m \u001b[38;5;28mopen\u001b[39m(\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mmminference_best_recalls.json\u001b[39m\u001b[38;5;124m\"\u001b[39m, \u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mr\u001b[39m\u001b[38;5;124m\"\u001b[39m) \u001b[38;5;28;01mas\u001b[39;00m f:\n",559      "File \u001b[0;32m~/miniconda3/envs/llava/lib/python3.10/site-packages/IPython/core/interactiveshell.py:324\u001b[0m, in \u001b[0;36m_modified_open\u001b[0;34m(file, *args, **kwargs)\u001b[0m\n\u001b[1;32m    317\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m file \u001b[38;5;129;01min\u001b[39;00m {\u001b[38;5;241m0\u001b[39m, \u001b[38;5;241m1\u001b[39m, \u001b[38;5;241m2\u001b[39m}:\n\u001b[1;32m    318\u001b[0m     \u001b[38;5;28;01mraise\u001b[39;00m \u001b[38;5;167;01mValueError\u001b[39;00m(\n\u001b[1;32m    319\u001b[0m         \u001b[38;5;124mf\u001b[39m\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mIPython won\u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mt let you open fd=\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mfile\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m by default \u001b[39m\u001b[38;5;124m\"\u001b[39m\n\u001b[1;32m    320\u001b[0m         \u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mas it is likely to crash IPython. If you know what you are doing, \u001b[39m\u001b[38;5;124m\"\u001b[39m\n\u001b[1;32m    321\u001b[0m         \u001b[38;5;124m\"\u001b[39m\u001b[38;5;124myou can use builtins\u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124m open.\u001b[39m\u001b[38;5;124m\"\u001b[39m\n\u001b[1;32m    322\u001b[0m     )\n\u001b[0;32m--> 324\u001b[0m \u001b[38;5;28;01mreturn\u001b[39;00m \u001b[43mio_open\u001b[49m\u001b[43m(\u001b[49m\u001b[43mfile\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;241;43m*\u001b[39;49m\u001b[43margs\u001b[49m\u001b[43m,\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;241;43m*\u001b[39;49m\u001b[38;5;241;43m*\u001b[39;49m\u001b[43mkwargs\u001b[49m\u001b[43m)\u001b[49m\n",560      "\u001b[0;31mFileNotFoundError\u001b[0m: [Errno 2] No such file or directory: 'mminference_best_patterns.json'"561     ]562    }563   ],564   "source": [565    "import json\n",566    "from collections import Counter\n",567    "\n",568    "seq_len = 20020\n",569    "\n",570    "with open(\"mminference_best_patterns.json\", \"r\") as f:\n",571    "    data = json.load(f)\n",572    "\n",573    "with open(\"mminference_best_recalls.json\", \"r\") as f:\n",574    "    recalls = json.load(f)\n",575    "\n",576    "pattern_counter = Counter()\n",577    "all_pattern_counter = Counter()\n",578    "\n",579    "for layer, heads in data.items():\n",580    "    print(\"=\" * 10, f\"{layer}\")\n",581    "    for head, pattern in heads.items():\n",582    "        recall = recalls[layer][head]\n",583    "\n",584    "        pattern_counter[pattern[0]] += 1\n",585    "        all_pattern_counter[pattern[0]] += 1\n",586    "\n",587    "        if pattern[0] == \"grid_attn\":\n",588    "            pass\n",589    "            # print(f\"head {head}\")\n",590    "            # print(f\"recall:\")\n",591    "            # print(json.dumps(recall, indent=1))\n",592    "            # print(f\"pattern:\")\n",593    "            # print(json.dumps(pattern))\n",594    "            # print('-' * 10)\n",595    "\n",596    "    print(f\"pattern_counter: {pattern_counter}\")\n",597    "    pattern_counter.clear()\n",598    "\n",599    "all_pattern_counter\n",600    "\n",601    "    "602   ]603  },604  {605   "cell_type": "code",606   "execution_count": 10,607   "metadata": {},608   "outputs": [609    {610     "data": {611      "image/png": "iVBORw0KGgoAAAANSUhEUgAAAv8AAAL7CAYAAABqauo0AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjAsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvlHJYcgAAAAlwSFlzAAASdAAAEnQB3mYfeAAAFk5JREFUeJzt3UGK41qAptFwIwy1hfZaavW5D8+cI88TbMPtQUNRUVbUs/KFQra/c4aBTNzhx4Vf2o0xxgcAAPD2/s/WBwAAAH6G+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR09IfXM/Hu7/9x//9z285DAAAsMztcnr4WTf/AAAQIf4BACBC/AMAQIT4BwCAiMWD3zl/fv+6+5sRMAAAPBc3/wAAECH+AQAgQvwDAECE+AcAgIhvGfzOmRsBf3wYAgMAwFbc/AMAQIT4BwCACPEPAAAR4h8AACJWG/x+xdeAAQBgG27+AQAgQvwDAECE+AcAgAjxDwAAET8++J1jBAwAAOtz8w8AABHiHwAAIsQ/AABEiH8AAIh4isHvHCNgAAD4Xm7+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAinvZtP3O8AQgAAP6em38AAIgQ/wAAECH+AQAgQvwDAEDESw1+5xgBAwDAY9z8AwBAhPgHAIAI8Q8AABHiHwAAIl5+8DvHCBgAAO65+QcAgAjxDwAAEeIfAAAixD8AAES85eB3ztwI+OPDEBgAgA43/wAAECH+AQAgQvwDAECE+AcAgIjM4PcrvgYMAECFm38AAIgQ/wAAECH+AQAgQvwDAEBEfvA7xwgYAIB35OYfAAAixD8AAESIfwAAiBD/AAAQYfD7ICNgAABenZt/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAIb/v5F7wBCACAV+LmHwAAIsQ/AABEiH8AAIgQ/wAAEGHw+82MgAEAeFZu/gEAIEL8AwBAhPgHAIAI8Q8AABEGvz9gbgT88WEIDADAz3LzDwAAEeIfAAAixD8AAESIfwAAiDD43ZCvAQMA8JPc/AMAQIT4BwCACPEPAAAR4h8AACJ2Y4yx5AfT/rDWWXK++vLv/2QEDADAV26X08PPuvkHAIAI8Q8AABHiHwAAIsQ/AABE+MLvC/AlYAAAvoObfwAAiBD/AAAQIf4BACBC/AMAQITB74syAgYAYCk3/wAAECH+AQAgQvwDAECE+AcAgIjdGGMs+cH1fFzrLKzACBgA4L3dLqeHn3XzDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQsfhtP9P+sNZZcv78/rXZ//YWIACA9+BtPwAAwB3xDwAAEeIfAAAixD8AAERMWx+AbcyNjY2AAQDem5t/AACIEP8AABAh/gEAIEL8AwBAhMEv/8UIGADgvbn5BwCACPEPAAAR4h8AACLEPwAAROzGGGPJD67n41pn4UUYAQMAPI/b5fTws27+AQAgQvwDAECE+AcAgAjxDwAAEYsHv9P+sNZZcua+qPuqjIABALZh8AsAANwR/wAAECH+AQAgQvwDAEDEtPUBeA9z42UjYACA5+LmHwAAIsQ/AABEiH8AAIgQ/wAAEGHwy2q++oKxITAAwDbc/AMAQIT4BwCACPEPAAAR4h8AACLEPwAAROzGGGPJD67n41pnIcwbgAAA/s7tcnr4WTf/AAAQIf4BACBC/AMAQIT4BwCAiMWD32l/WOssOX9+/9r6CE/NCBgA4J8Z/AIAAHfEPwAARIh/AACIEP8AABAxbX0A+MrcINoIGADg77n5BwCACPEPAAAR4h8AACLEPwAARBj88lKMgAEA/p6bfwAAiBD/AAAQIf4BACBC/AMAQMRujDGW/OB6Pq51Fvg2RsAAQMXtcnr4WTf/AAAQIf4BACBC/AMAQIT4BwCAiMWD32l/WOssOXNfq2VdhsAAwLsx+AUAAO6IfwAAiBD/AAAQIf4BACBi2voA8JPmRtZGwABAhZt/AACIEP8AABAh/gEAIEL8AwBAhMEveUbAAECFm38AAIgQ/wAAECH+AQAgQvwDAECE+AcAgIjdGGMs+cH1fFzrLPDUvAEIAHhGt8vp4Wfd/AMAQIT4BwCACPEPAAAR4h8AACIWD36n/WGts+T8+f1r6yPwLxkBAwBbM/gFAADuiH8AAIgQ/wAAECH+AQAgYtr6APDK5kbbRsAAwLNy8w8AABHiHwAAIsQ/AABEiH8AAIgw+IVv9tWXmw2BAYCtufkHAIAI8Q8AABHiHwAAIsQ/AABE7MYYY8kPrufjWmeBHCNgAODful1ODz/r5h8AACLEPwAARIh/AACIEP8AABCxePA77Q9rnSXnqy/B0mYEDAAsYfALAADcEf8AABAh/gEAIEL8AwBAxLT1AYDP5obgRsAAwHdw8w8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAEOFtP/ACvAEIAPgObv4BACBC/AMAQIT4BwCACPEPAAARuzHGWPKD6/m41lmAf8kIGAB6bpfTw8+6+QcAgAjxDwAAEeIfAAAixD8AAEQsHvxO+8NaZ8mZ+2orrMEQGADel8EvAABwR/wDAECE+AcAgAjxDwAAEdPWBwDWNzcuNwIGgB43/wAAECH+AQAgQvwDAECE+AcAgAiDX4gyAgaAHjf/AAAQIf4BACBC/AMAQIT4BwCAiN0YYyz5wfV8XOsswBMyAgaA53a7nB5+1s0/AABEiH8AAIgQ/wAAECH+AQAgYvHgd9of1jpLztwXVuEVGAEDwPMw+AUAAO6IfwAAiBD/AAAQIf4BACBi2voAwOuZG6sbAQPA83PzDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQ4W0/wLeYewPQx4e3AAHAM3HzDwAAEeIfAAAixD8AAESIfwAAiNiNMcaSH1zPx7XOAkQYAQPA97ldTg8/6+YfAAAixD8AAESIfwAAiBD/AAAQsXjwO+0Pa50l56svokKRETAA/B2DXwAA4I74BwCACPEPAAAR4h8AACKmrQ8A8PExP4A3AgaA7+XmHwAAIsQ/AABEiH8AAIgQ/wAAEGHwCzwtI2AA+F5u/gEAIEL8AwBAhPgHAIAI8Q8AABG7McZY8oPr+bjWWQD+ihEwAGW3y+nhZ938AwBAhPgHAIAI8Q8AABHiHwAAIhYPfqf9Ya2z5Mx9vRT4PobAABQY/AIAAHfEPwAARIh/AACIEP8AABAh/gEAIGLa+gAAa5l7o5Y3AAFQ5uYfAAAixD8AAESIfwAAiBD/AAAQYfALpBgBA1Dm5h8AACLEPwAARIh/AACIEP8AABCxG2OMJT+4no9rnQXgaRgBA/AqbpfTw8+6+QcAgAjxDwAAEeIfAAAixD8AAEQsHvxO+8NaZ8mZ+9Io8LyMgAF4Rga/AADAHfEPAAAR4h8AACLEPwAARExbHwDgVcyN9I2AAXglbv4BACBC/AMAQIT4BwCACPEPAAARBr8A/8JXX+o2BAbgGbn5BwCACPEPAAAR4h8AACLEPwAAROzGGGPJD67n41pnAXhrRsAArOF2OT38rJt/AACIEP8AABAh/gEAIEL8AwBAxOLB77Q/rHUWgLfx1Zd//ycjYAD+LYNfAADgjvgHAIAI8Q8AABHiHwAAIsQ/AABETFsfAKBs7q1A3gAEwFrc/AMAQIT4BwCACPEPAAAR4h8AACIMfgGejBEwAGtx8w8AABHiHwAAIsQ/AABEiH8AAIjYjTHGkh9cz8e1zgLAAkbAAHx8fHzcLqeHn3XzDwAAEeIfAAAixD8AAESIfwAAiFg8+J32h7XOAvA25r7S+1MMgQFaDH4BAIA74h8AACLEPwAARIh/AACImLY+AADfa25sbAQMwMeHm38AAMgQ/wAAECH+AQAgQvwDAECEwS9AgBEwAB8fbv4BACBD/AMAQIT4BwCACPEPAAARuzHGWPKD6/m41lkA2JgRMMDruV1ODz/r5h8AACLEPwAARIh/AACIEP8AABAh/gEAIGLx236m/WGtswC8jT+/f219hG/jDUAAz83bfgAAgDviHwAAIsQ/AABEiH8AAIiYtj4AAM9tbrxsBAzwmtz8AwBAhPgHAIAI8Q8AABHiHwAAIgx+AVjsqy8YGwIDPDc3/wAAECH+AQAgQvwDAECE+AcAgIjdGGMs+cH1fFzrLAC8ISNggHXdLqeHn3XzDwAAEeIfAAAixD8AAESIfwAAiFg8+J32h7XOAvA2vvoCLv+fETDA9zH4BQAA7oh/AACIEP8AABAh/gEAIGLa+gAA9MwNoo2AAdbn5h8AACLEPwAARIh/AACIEP8AABBh8AvAUzACBlifm38AAIgQ/wAAECH+AQAgQvwDAEDEbowxlvzgej6udRYA+EdGwACf3S6nh5918w8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAELH4bT/T/rDWWQDgkz+/fz38rLcAAVXe9gMAANwR/wAAECH+AQAgQvwDAEDEtPUBAOA7zI2DjYABPnPzDwAAEeIfAAAixD8AAESIfwAAiDD4BeBtGQEDfObmHwAAIsQ/AABEiH8AAIgQ/wAAELEbY4wlP7iej2udBQA2YQQMvLLb5fTws27+AQAgQvwDAECE+AcAgAjxDwAAEYsHv9P+sNZZAOCTuS/0/hQjYOBVGPwCAAB3xD8AAESIfwAAiBD/AAAQMW19AAB4RnNjYyNg4NW5+QcAgAjxDwAAEeIfAAAixD8AAEQY/ALAg7764rAhMPAq3PwDAECE+AcAgAjxDwAAEeIfAAAixD8AAETsxhhjyQ+u5+NaZwGAt+ENQMBPuV1ODz/r5h8AACLEPwAARIh/AACIEP8AABCxePA77Q9rnQUAPvnz+9fWR/hWRsDAGgx+AQCAO+IfAAAixD8AAESIfwAAiJi2PgAAVMwNmI2AgZ/k5h8AACLEPwAARIh/AACIEP8AABBh8AsAGzICBn6Sm38AAIgQ/wAAECH+AQAgQvwDAEDEbowxlvzgej6udRYA4AtGwMBXbpfTw8+6+QcAgAjxDwAAEeIfAAAixD8AAEQsHvxO+8NaZwGAT+a+fstnhsCAwS8AAHBH/AMAQIT4BwCACPEPAAAR09YHAAD+3two2ggY+IqbfwAAiBD/AAAQIf4BACBC/AMAQITBLwC8GSNg4Ctu/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAInZjjLHkB9fzca2zAAA/yBuA4D3cLqeHn3XzDwAAEeIfAAAixD8AAESIfwAAiFg8+J32h7XOAgCf/Pn9a+sj5BgBw+sx+AUAAO6IfwAAiBD/AAAQIf4BACBi2voAAMDz+GpkbQgM78HNPwAARIh/AACIEP8AABAh/gEAIMLgFwD4R3NDYCNgeD1u/gEAIEL8AwBAhPgHAIAI8Q8AABG7McZY8oPr+bjWWQCAF2cEDD/vdjk9/KybfwAAiBD/AAAQIf4BACBC/AMAQMTiwe+0P6x1FgD4ZO6rsrweI2BYl8EvAABwR/wDAECE+AcAgAjxDwAAEdPWBwAA3tvccNsIGLbh5h8AACLEPwAARIh/AACIEP8AABAh/gEAIMLbfgCAH+cNQLANN/8AABAh/gEAIEL8AwBAhPgHAICI3RhjLPnB9Xxc6ywAAJ8YAcM/u11ODz/r5h8AACLEPwAARIh/AACIEP8AABCxePA77Q9rnQUAPpn7Cix8fBgCw39n8AsAANwR/wAAECH+AQAgQvwDAEDEtPUBAACWmhuDGwHDP3PzDwAAEeIfAAAixD8AAESIfwAAiDD4BQDeghEw/DM3/wAAECH+AQAgQvwDAECE+AcAgIjdGGMs+cH1fFzrLAAAqzMC5t3cLqeHn3XzDwAAEeIfAAAixD8AAESIfwAAiFg8+J32h7XOAgCfzH2xFdZgBMwrM/gFAADuiH8AAIgQ/wAAECH+AQAgYtr6AAAAW5sblxsB847c/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARHjbDwDAjLk3AH18eAsQr83NPwAARIh/AACIEP8AABAh/gEAIGI3xhhLfnA9H9c6CwDASzICZku3y+nhZ938AwBAhPgHAIAI8Q8AABHiHwAAIhYPfqf9Ya2zAMAnX31hFV6BETA/xeAXAAC4I/4BACBC/AMAQIT4BwCAiGnrAwAAvKO5wboRMFtz8w8AABHiHwAAIsQ/AABEiH8AAIgw+AUA+CFGwGzNzT8AAESIfwAAiBD/AAAQIf4BACBiN8YYS35wPR/XOgsAAB9GwCxzu5weftbNPwAARIh/AACIEP8AABAh/gEAIGLx4HfaH9Y6CwB8Mvc1VCgzBGaOwS8AAHBH/AMAQIT4BwCACPEPAAAR4h8AACKmrQ8AAMBj5t6A5Q1ALOHmHwAAIsQ/AABEiH8AAIgQ/wAAEGHwCwDwwoyAWcLNPwAARIh/AACIEP8AABAh/gEAIGI3xhhLfnA9H9c6CwAAKzECfl+3y+nhZ938AwBAhPgHAIAI8Q8AABHiHwAAIhYPfqf9Ya2zAMAnc18uBb6PEfB7MPgFAADuiH8AAIgQ/wAAECH+AQAgYtr6AAAAbGNuVG8E/N7c/AMAQIT4BwCACPEPAAAR4h8AACIMfgEA+C9ffVnbEPg9uPkHAIAI8Q8AABHiHwAAIsQ/AABE7MYYY8kPrufjWmcBAOCFGAE/h9vl9PCzbv4BACBC/AMAQIT4BwCACPEPAAARiwe/0/6w1lkA4JOvvjQKPC8j4J9n8AsAANwR/wAAECH+AQAgQvwDAECE+AcAgIhp6wMAAPA+5t7S5Q1Az8PNPwAARIh/AACIEP8AABAh/gEAIMLgFwCAVRkBPw83/wAAECH+AQAgQvwDAECE+AcAgIjdGGMs+cH1fFzrLAAAhBkB/53b5fTws27+AQAgQvwDAECE+AcAgAjxDwAAEYsHv9P+sNZZAACImPvq71cMgf93Br8AAMAd8Q8AABHiHwAAIsQ/AABETFsfAAAA/jdz42Aj4L/j5h8AACLEPwAARIh/AACIEP8AABBh8AsAwMsxAv47bv4BACBC/AMAQIT4BwCACPEPAAARuzHGWPKD6/m41lkAAOBbFUbAt8vp4Wfd/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARCx+28+0P6x1FgAAIv78/rXZ/363NwB52w8AAHBH/AMAQIT4BwCACPEPAAAR09YHAACAnzQ3Nn63EfBX3PwDAECE+AcAgAjxDwAAEeIfAAAiDH4BAMj76ovD7zYEdvMPAAAR4h8AACLEPwAARIh/AACI2I0xxpIfXM/Htc4CAABP79lGwLfL6eFn3fwDAECE+AcAgAjxDwAAEeIfAAAiFg9+p/1hrbMAABDx1Rd1X9WWI2CDXwAA4I74BwCACPEPAAAR4h8AACKmrQ8AAACvbm7A/GxfAv74cPMPAAAZ4h8AACLEPwAARIh/AACIMPgFAIAVPOMI2M0/AABEiH8AAIgQ/wAAECH+AQAgYjfGGEt+cD0f1zoLAADk/NsR8O1yevhZN/8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABGL3/Yz7Q9rnQUAgIg/v39tfYSn9+hbgLztBwAAuCP+AQAgQvwDAECE+AcAgIhp6wMAAAD35kbRj46Av+LmHwAAIsQ/AABEiH8AAIgQ/wAAEGHwCwAAL+LffhnZzT8AAESIfwAAiBD/AAAQIf4BACBiN8YYWx8CAABYn5t/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAg4v8BdaB5IUAUktoAAAAASUVORK5CYII=",612      "text/plain": [613       "<Figure size 960x960 with 1 Axes>"614      ]615     },616     "metadata": {},617     "output_type": "display_data"618    }619   ],620   "source": [621    "import torch\n",622    "import seaborn as sns\n",623    "import matplotlib.pyplot as plt\n",624    "\n",625    "num_tokens = 128+16\n",626    "\n",627    "def plot_mask(mask):\n",628    "    plt.figure(figsize=(8, 8), dpi=120)\n",629    "    sns.heatmap(mask.numpy(), cbar=False)\n",630    "    plt.axis('off')\n",631    "    plt.show()\n",632    "\n",633    "modality_boundaries = torch.tensor([0, 16, 64, 80, 128, 144])\n",634    "mask = torch.zeros((num_tokens, num_tokens), dtype=torch.int32)\n",635    "\n",636    "for i in range(modality_boundaries.shape[0]):\n",637    "    prev_boundary = modality_boundaries[i-1] if i > 0 else 0\n",638    "    for j in range(prev_boundary, modality_boundaries[i]):\n",639    "        for k in range(prev_boundary, modality_boundaries[i]):\n",640    "            if j >= k:\n",641    "                mask[j, k] = 1\n",642    "                if j % 4 == 0:\n",643    "                    mask[j, :k] = 1\n",644    "plot_mask(mask)\n",645    "\n"646   ]647  },648  {649   "cell_type": "code",650   "execution_count": null,651   "metadata": {},652   "outputs": [],653   "source": [654    "shuffle_index = torch.arange(num_tokens, dtype=torch.int64).reshape((-1, stride)).T.flatten()"655   ]656  },657  {658   "cell_type": "code",659   "execution_count": 11,660   "metadata": {},661   "outputs": [662    {663     "name": "stdout",664     "output_type": "stream",665     "text": [666      "[0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 64, 65, 66, 67, 68, 69, 70, 71, 72, 73, 74, 75, 76, 77, 78, 79, 80, 128, 129, 130, 131, 132, 133, 134, 135, 136, 137, 138, 139, 140, 141, 142, 143, 144, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 32, 33, 34, 35, 36, 37, 38, 39, 40, 41, 42, 43, 44, 45, 46, 47, 48, 49, 50, 51, 52, 53, 54, 55, 56, 57, 58, 59, 60, 61, 62, 63, 64, 80, 81, 82, 83, 84, 85, 86, 87, 88, 89, 90, 91, 92, 93, 94, 95, 96, 97, 98, 99, 100, 101, 102, 103, 104, 105, 106, 107, 108, 109, 110, 111, 112, 113, 114, 115, 116, 117, 118, 119, 120, 121, 122, 123, 124, 125, 126, 127, 128]\n"667     ]668    }669   ],670   "source": [671    "def create_typed_spans(boundaries, types):\n",672    "    result = []\n",673    "    for i in range(len(types)):\n",674    "        if types[i] == 't':\n",675    "            start, end = boundaries[i], boundaries[i+1]\n",676    "            result.extend(range(start, end + 1))\n",677    "    for i in range(len(types)):\n",678    "        if types[i] == 'v':\n",679    "            start, end = boundaries[i], boundaries[i+1]\n",680    "            result.extend(range(start, end + 1))\n",681    "    return result\n",682    "    \n",683    "span_boundaries = (0, 16, 64, 80, 128, 144)\n",684    "span_types = ('t', 'v', 't', 'v', 't')\n",685    "\n",686    "result = create_typed_spans(span_boundaries, span_types)\n",687    "print(result)"688   ]689  },690  {691   "cell_type": "code",692   "execution_count": 13,693   "metadata": {},694   "outputs": [695    {696     "name": "stdout",697     "output_type": "stream",698     "text": [699      "tensor([ 16,  17,  18,  19,  20,  21,  22,  23,  24,  25,  26,  27,  28,  29,\n",700      "         30,  31,  32,  33,  34,  35,  36,  37,  38,  39,  40,  41,  42,  43,\n",701      "         44,  45,  46,  47,  48,  49,  50,  51,  52,  53,  54,  55,  56,  57,\n",702      "         58,  59,  60,  61,  62,  63,  80,  81,  82,  83,  84,  85,  86,  87,\n",703      "         88,  89,  90,  91,  92,  93,  94,  95,  96,  97,  98,  99, 100, 101,\n",704      "        102, 103, 104, 105, 106, 107, 108, 109, 110, 111, 112, 113, 114, 115,\n",705      "        116, 117, 118, 119, 120, 121, 122, 123, 124, 125, 126, 127,   0,   1,\n",706      "          2,   3,   4,   5,   6,   7,   8,   9,  10,  11,  12,  13,  14,  15,\n",707      "         64,  65,  66,  67,  68,  69,  70,  71,  72,  73,  74,  75,  76,  77,\n",708      "         78,  79, 128, 129, 130, 131, 132, 133, 134, 135, 136, 137, 138, 139,\n",709      "        140, 141, 142, 143])\n"710     ]711    }712   ],713   "source": [714    "import numpy as np\n",715    "\n",716    "def create_typed_spans(boundaries, types):\n",717    "    # Convert inputs to numpy arrays if they aren't already\n",718    "    boundaries = np.asarray(boundaries)\n",719    "    \n",720    "    # Create masks for t and v types\n",721    "    t_mask = np.array([t == 't' for t in types])\n",722    "    v_mask = np.array([t == 'v' for t in types])\n",723    "    \n",724    "    # Get the start and end points for each type\n",725    "    t_starts = boundaries[:-1][t_mask]\n",726    "    t_ends = boundaries[1:][t_mask]\n",727    "    v_starts = boundaries[:-1][v_mask]\n",728    "    v_ends = boundaries[1:][v_mask]\n",729    "    \n",730    "    # Create arrays for each type using np.arange\n",731    "    t_arrays = [np.arange(start, end) for start, end in zip(t_starts, t_ends)]\n",732    "    v_arrays = [np.arange(start, end) for start, end in zip(v_starts, v_ends)]\n",733    "    \n",734    "    # Concatenate all arrays\n",735    "    # result = np.concatenate(t_arrays + v_arrays)\n",736    "    result = np.concatenate(v_arrays+t_arrays)\n",737    "    \n",738    "    return torch.tensor(result)\n",739    "\n",740    "span_boundaries = (0, 16, 64, 80, 128, 144)\n",741    "span_types = ('t', 'v', 't', 'v', 't')\n",742    "\n",743    "result = create_typed_spans(span_boundaries, span_types)\n",744    "print(result)"745   ]746  },747  {748   "cell_type": "code",749   "execution_count": 15,750   "metadata": {},751   "outputs": [752    {753     "data": {754      "image/png": "iVBORw0KGgoAAAANSUhEUgAAAv8AAAL7CAYAAABqauo0AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjAsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvlHJYcgAAAAlwSFlzAAASdAAAEnQB3mYfeAAAFlBJREFUeJzt3UGO4lqAptGgZCH1Foq11OpzH8zIEfOUAOnWoFvVCuGoh/OFw8B3zjCFI+/w05V+ezfGGB8AAMDb+4+tDwAAAPwM8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAipqUPXM/HNc7B//N//vO/tj4CAAAv5HY5PfxbN/8AABAh/gEAIEL8AwBAhPgHAICI3RhjLHlg2h/WOkvOn9+/HvqdETAAAF8x+AUAAO6IfwAAiBD/AAAQIf4BACBi8Rd++Xlzw2AjYAAAlnLzDwAAEeIfAAAixD8AAESIfwAAiDD4fVFGwAAALOXmHwAAIsQ/AABEiH8AAIgQ/wAAELEbY4wlD1zPx7XOwgqMgAEA3tvtcnr4t27+AQAgQvwDAECE+AcAgAjxDwAAEYsHv9P+sNZZcua+0vtTDIEBAN6DwS8AAHBH/AMAQIT4BwCACPEPAAAR4h8AACKmrQ/ANubeNOQNQAAA783NPwAARIh/AACIEP8AABAh/gEAIMLgl/9hBAwA8N7c/AMAQIT4BwCACPEPAAAR4h8AACJ2Y4yx5IHr+bjWWXgRRsAAAM/jdjk9/Fs3/wAAECH+AQAgQvwDAECE+AcAgIjFg99pf1jrLDlzX9R9VUbAAADbMPgFAADuiH8AAIgQ/wAAECH+AQAgYtr6ALyHufGyETAAwHNx8w8AABHiHwAAIsQ/AABEiH8AAIgw+GU1X33B2BAYAGAbbv4BACBC/AMAQIT4BwCACPEPAAARuzHGWPLA9Xxc6yyEGQEDAPyd2+X08G/d/AMAQIT4BwCACPEPAAAR4h8AACLEPwAARCx+28+0P6x1lpw/v39tfYSn5g1AAAD/zNt+AACAO+IfAAAixD8AAESIfwAAiJi2PgB8ZW4QbQQMAPD33PwDAECE+AcAgAjxDwAAEeIfAAAiDH55KUbAAAB/z80/AABEiH8AAIgQ/wAAECH+AQAgYjfGGEseuJ6Pa50Fvo0RMABQcbucHv6tm38AAIgQ/wAAECH+AQAgQvwDAEDE4sHvtD+sdZacua/Vsi5DYADg3Rj8AgAAd8Q/AABEiH8AAIgQ/wAAEDFtfQD4SXMjayNgAKDCzT8AAESIfwAAiBD/AAAQIf4BACDC4Jc8I2AAoMLNPwAARIh/AACIEP8AABAh/gEAIGI3xhhLHriej2udBZ6aETAA8Ixul9PDv3XzDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQsfhtP9P+sNZZcv78/rX1EfiXvAEIANiat/0AAAB3xD8AAESIfwAAiBD/AAAQMW19AHhlc6NtI2AA4Fm5+QcAgAjxDwAAEeIfAAAixD8AAEQY/MI3++rLzYbAAMDW3PwDAECE+AcAgAjxDwAAEeIfAAAidmOMseSB6/m41lkgxwgYAPi3bpfTw7918w8AABHiHwAAIsQ/AABEiH8AAIhYPPid9oe1zpLz1ZdgaTMCBgCWMPgFAADuiH8AAIgQ/wAAECH+AQAgYtr6AMBnc0NwI2AA4Du4+QcAgAjxDwAAEeIfAAAixD8AAEQY/MILMAIGAL6Dm38AAIgQ/wAAECH+AQAgQvwDAECE+AcAgIjdGGMseeB6Pq51FuBf8gYgAOi5XU4P/9bNPwAARIh/AACIEP8AABAh/gEAIGLx4HfaH9Y6S86f37+2PgIRhsAA8L4MfgEAgDviHwAAIsQ/AABEiH8AAIiYtj4AsL65cbkRMAD0uPkHAIAI8Q8AABHiHwAAIsQ/AABEGPxClBEwAPS4+QcAgAjxDwAAEeIfAAAixD8AAETsxhhjyQPX83GtswBPyAgYAJ7b7XJ6+Ldu/gEAIEL8AwBAhPgHAIAI8Q8AABGLB7/T/rDWWXLmvrAKr8AIGACeh8EvAABwR/wDAECE+AcAgAjxDwAAEdPWBwBez9xY3QgYAJ6fm38AAIgQ/wAAECH+AQAgQvwDAECEwS/wLb76YrUhMAA8Dzf/AAAQIf4BACBC/AMAQIT4BwCACPEPAAARuzHGWPLA9Xxc6yxAhDcAAcD3uV1OD//WzT8AAESIfwAAiBD/AAAQIf4BACBi8eB32h/WOkvOn9+/tj4CPA0jYAD4Owa/AADAHfEPAAAR4h8AACLEPwAARExbHwDg42N+AG8EDADfy80/AABEiH8AAIgQ/wAAECH+AQAgwuAXeFpGwADwvdz8AwBAhPgHAIAI8Q8AABHiHwAAInZjjLHkgev5uNZZAP6KETAAZbfL6eHfuvkHAIAI8Q8AABHiHwAAIsQ/AABELB78TvvDWmfJmft6KfB9DIEBKDD4BQAA7oh/AACIEP8AABAh/gEAIGLa+gAAa5kb1RsBA1Dm5h8AACLEPwAARIh/AACIEP8AABAh/gEAIMLbfoAUbwACoMzNPwAARIh/AACIEP8AABAh/gEAIGI3xhhLHriej2udBeBpGAED8Cpul9PDv3XzDwAAEeIfAAAixD8AAESIfwAAiFg8+J32h7XOkjP3pVHgeRkBA/CMDH4BAIA74h8AACLEPwAARIh/AACImLY+AMCrmBvpGwED8Erc/AMAQIT4BwCACPEPAAAR4h8AACIMfgH+ha++1G0IDMAzcvMPAAAR4h8AACLEPwAARIh/AACI2I0xxpIHrufjWmcBgH9kTA3w2e1yevi3bv4BACBC/AMAQIT4BwCACPEPAAARiwe/0/6w1lkA4JOvvqA8xxAYqDL4BQAA7oh/AACIEP8AABAh/gEAIGLa+gAA8B3mxsFGwACfufkHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIjwth8A3pY3AAF85uYfAAAixD8AAESIfwAAiBD/AAAQsRtjjCUPXM/Htc4CAJswAgZe2e1yevi3bv4BACBC/AMAQIT4BwCACPEPAAARiwe/0/6w1lkA4JO5L/T+FCNg4FUY/AIAAHfEPwAARIh/AACIEP8AABAxbX0AAHhGc2NjI2Dg1bn5BwCACPEPAAAR4h8AACLEPwAARBj8AsCDvvrisCEw8Crc/AMAQIT4BwCACPEPAAAR4h8AACJ2Y4yx5IHr+bjWWQDgbRgBAz/ldjk9/Fs3/wAAECH+AQAgQvwDAECE+AcAgIjFg99pf1jrLADwyVdf1H1VRsDAGgx+AQCAO+IfAAAixD8AAESIfwAAiBD/AAAQMW19AAComHt7kTcAAT/JzT8AAESIfwAAiBD/AAAQIf4BACDC4BcANmQEDPwkN/8AABAh/gEAIEL8AwBAhPgHAICI3RhjLHngej6udRYA4AtGwMBXbpfTw7918w8AABHiHwAAIsQ/AABEiH8AAIhYPPid9oe1zgIAn8x9/ZbPDIEBg18AAOCO+AcAgAjxDwAAEeIfAAAipq0PAAD8vblRtBEw8BU3/wAAECH+AQAgQvwDAECE+AcAgAiDXwB4M0bAwFfc/AMAQIT4BwCACPEPAAAR4h8AACJ2Y4yx5IHr+bjWWQCAH2QEDO/hdjk9/Fs3/wAAECH+AQAgQvwDAECE+AcAgIjFg99pf1jrLADwydyXalmXETC8HoNfAADgjvgHAIAI8Q8AABHiHwAAIsQ/AABETFsfAAB4Hl+9YclbgOA9uPkHAIAI8Q8AABHiHwAAIsQ/AABEGPwCAP9obghsBAyvx80/AABEiH8AAIgQ/wAAECH+AQAgYjfGGEseuJ6Pa50FAHhxRsDw826X08O/dfMPAAAR4h8AACLEPwAARIh/AACIWDz4nfaHtc4CAJ/MfVWW12MEDOsy+AUAAO6IfwAAiBD/AAAQIf4BACBi2voAAMB7mxtuGwHDNtz8AwBAhPgHAIAI8Q8AABHiHwAAIgx+AYAfZwQM23DzDwAAEeIfAAAixD8AAESIfwAAiNiNMcaSB67n41pnAQD4xAgY/tntcnr4t27+AQAgQvwDAECE+AcAgAjxDwAAEeIfAAAiFr/tZ9of1joLAHzy5/evrY/Ak/IWIPj/vO0HAAC4I/4BACBC/AMAQIT4BwCAiGnrAwAALDU3BjcChn/m5h8AACLEPwAARIh/AACIEP8AABBh8AsAvAUjYPhnbv4BACBC/AMAQIT4BwCACPEPAAARuzHGWPLA9Xxc6ywAAKszAubd3C6nh3/r5h8AACLEPwAARIh/AACIEP8AABCxePA77Q9rnQUAPpn7YiuswQiYV2bwCwAA3BH/AAAQIf4BACBC/AMAQMS09QEAALY2Ny43AuYdufkHAIAI8Q8AABHiHwAAIsQ/AABEGPwCAMz46gvThsC8Mjf/AAAQIf4BACBC/AMAQIT4BwCAiN0YYyx54Ho+rnUWAICXZATMlm6X08O/dfMPAAAR4h8AACLEPwAARIh/AACIEP8AABCx+G0/0/6w1lkA4JM/v39tfQT4a94AxE/xth8AAOCO+AcAgAjxDwAAEeIfAAAipq0PAADwjuYG60bAbM3NPwAARIh/AACIEP8AABAh/gEAIMLgFwDghxgBszU3/wAAECH+AQAgQvwDAECE+AcAgIjdGGMseeB6Pq51FgAAPoyAWeZ2OT38Wzf/AAAQIf4BACBC/AMAQIT4BwCAiMWD32l/WOssAPDJ3NdQocwQmDkGvwAAwB3xDwAAEeIfAAAixD8AAERMWx8AAIDHzI3gjYBZws0/AABEiH8AAIgQ/wAAECH+AQAgwuAXAOCFGQGzhJt/AACIEP8AABAh/gEAIEL8AwBAhPgHAICI3RhjLHngej6udRYAAFbiDUDv63Y5PfxbN/8AABAh/gEAIEL8AwBAhPgHAICIxYPfaX9Y6ywA8Mmf37+2PgK8NSPg92DwCwAA3BH/AAAQIf4BACBC/AMAQMS09QEAANjG3KjeCPi9ufkHAIAI8Q8AABHiHwAAIsQ/AABEGPwCAPA/vvqytiHwe3DzDwAAEeIfAAAixD8AAESIfwAAiNiNMcaSB67n41pnAQDghRgBP4fb5fTwb938AwBAhPgHAIAI8Q8AABHiHwAAIhYPfqf9Ya2zAMAnX31pFHheRsA/z+AXAAC4I/4BACBC/AMAQIT4BwCAiGnrAwAA8D7mhvpGwM/DzT8AAESIfwAAiBD/AAAQIf4BACDC4BcAgFUZAT8PN/8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABG7McZY8sD1fLz7N2ttAADYxu1yevi3bv4BACBC/AMAQIT4BwCACPEPAAAR03f8EZ9sBgCA5+fmHwAAIsQ/AABEiH8AAIgQ/wAAEPEtg985cyPgjw9DYAAA2IqbfwAAiBD/AAAQIf4BACBC/AMAQMRqg9+v+BowAABsw80/AABEiH8AAIgQ/wAAECH+AQAg4scHv3OMgAEAYH1u/gEAIEL8AwBAhPgHAIAI8Q8AABFPMfidYwQMAADfy80/AABEiH8AAIgQ/wAAECH+AQAg4mkHv3OMgAEA4O+5+QcAgAjxDwAAEeIfAAAixD8AAESIfwAAiHipt/3M8QYgAAB4jJt/AACIEP8AABAh/gEAIEL8AwBAxMsPfucYAQMAwD03/wAAECH+AQAgQvwDAECE+AcAgIi3HPzOmRsBf3wYAgMA0OHmHwAAIsQ/AABEiH8AAIgQ/wAAEJEZ/H7F14ABAKhw8w8AABHiHwAAIsQ/AABEiH8AAIjID37nGAEDAPCO3PwDAECE+AcAgAjxDwAAEeIfAAAiDH4fZAQMAMCrc/MPAAAR4h8AACLEPwAARIh/AACIMPj9F4yAAQB4JW7+AQAgQvwDAECE+AcAgAjxDwAAEQa/38wIGACAZ+XmHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgwtt+fsDcG4A+PrwFCACAn+XmHwAAIsQ/AABEiH8AAIgQ/wAAEGHwu6GvhsDAezLyB2Brbv4BACBC/AMAQIT4BwCACPEPAAARuzHGWPLAtD+sdRaAt/HooN8IGIB/63Y5PfxbN/8AABAh/gEAIEL8AwBAhPgHAIAIX/gF2NDcMNgIGIC1uPkHAIAI8Q8AABHiHwAAIsQ/AABEGPwCPBkjYADW4uYfAAAixD8AAESIfwAAiBD/AAAQsRtjjCUPXM/Htc4CwAJGwAB8fHx83C6nh3/r5h8AACLEPwAARIh/AACIEP8AABCxePA77Q9rnQXgbcx9pfenGAIDtBj8AgAAd8Q/AABEiH8AAIgQ/wAAECH+AQAgYtr6AAB8r7k3DXkDEAAfH27+AQAgQ/wDAECE+AcAgAjxDwAAEQa/AAFGwAB8fLj5BwCADPEPAAAR4h8AACLEPwAAROzGGGPJA9fzca2zALAxI2CA13O7nB7+rZt/AACIEP8AABAh/gEAIEL8AwBAxOLB77Q/rHUWgLcx90XdV2UEDPDcDH4BAIA74h8AACLEPwAARIh/AACImLY+AADPbW68bAQM8Jrc/AMAQIT4BwCACPEPAAAR4h8AACIMfgFY7KsvGBsCAzw3N/8AABAh/gEAIEL8AwBAhPgHAICI3RhjLHngej6udRYA3pARMMC6bpfTw7918w8AABHiHwAAIsQ/AABEiH8AAIhYPPid9oe1zgLwNr76Ai7/lxEwwPcx+AUAAO6IfwAAiBD/AAAQIf4BACBC/AMAQMS09QEA6Jl7G5I3AAGsz80/AABEiH8AAIgQ/wAAECH+AQAgwuAXgKdgBAywPjf/AAAQIf4BACBC/AMAQIT4BwCAiN0YYyx54Ho+rnUWAADCjPz/zu1yevi3bv4BACBC/AMAQIT4BwCACPEPAAARiwe/0/6w1lkAAIiY+6r3VwyB/3cGvwAAwB3xDwAAEeIfAAAixD8AAERMWx8AAAD+N3PjYCPgv+PmHwAAIsQ/AABEiH8AAIgQ/wAAEGHwCwDAyzEC/jtu/gEAIEL8AwBAhPgHAIAI8Q8AABG7McZY8sD1fFzrLAAA8K0KI+Db5fTwb938AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABELH7bz7Q/rHUWAAAi/vz+tdn//W5vAPK2HwAA4I74BwCACPEPAAAR4h8AACKmrQ8AAAA/aW5s/G4j4K+4+QcAgAjxDwAAEeIfAAAixD8AAEQY/AIAkPfVF4ffbQjs5h8AACLEPwAARIh/AACIEP8AABCxG2OMJQ9cz8e1zgIAAE/v2UbAt8vp4d+6+QcAgAjxDwAAEeIfAAAixD8AAEQsHvxO+8NaZwEAIOKrL+q+qi1HwAa/AADAHfEPAAAR4h8AACLEPwAARExbHwAAAF7d3ID52b4E/PHh5h8AADLEPwAARIh/AACIEP8AABBh8AsAACt4xhGwm38AAIgQ/wAAECH+AQAgQvwDAEDEbowxljxwPR/XOgsAAOT82xHw7XJ6+Ldu/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIha/7WfaH9Y6CwAAEX9+/9r6CE/v0bcAedsPAABwR/wDAECE+AcAgAjxDwAAEdPWBwAAAO7NjaIfHQF/xc0/AABEiH8AAIgQ/wAAECH+AQAgwuAXAABexL/9MrKbfwAAiBD/AAAQIf4BACBC/AMAQMRujDG2PgQAALA+N/8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAEDEfwPuvHkhbzICTgAAAABJRU5ErkJggg==",755      "text/plain": [756       "<Figure size 960x960 with 1 Axes>"757      ]758     },759     "metadata": {},760     "output_type": "display_data"761    },762    {763     "data": {764      "image/png": "iVBORw0KGgoAAAANSUhEUgAAAv8AAAL7CAYAAABqauo0AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjAsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvlHJYcgAAAAlwSFlzAAASdAAAEnQB3mYfeAAAFmVJREFUeJzt3UGK41iXgFG7EQG9hY619OprH565Rp7/YBtez6pJQq60Iq2Q7e+cYSFRb5KOjwv3aT/GGDsAAODt/dfWBwAAAH6G+AcAgAjxDwAAEeIfAAAixD8AAESIfwAAiBD/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR09IXLqfDl//23//zvw85DACwjv/8/dfWR3gbuodncz0f737W5B8AACLEPwAARIh/AACIEP8AABCxH2OMJS/MLfzOsQwDAADrs/ALAAB8If4BACBC/AMAQIT4BwCAiMVf+L3XrS8JWgQGAIBtmPwDAECE+AcAgAjxDwAAEeIfAAAiVlv4vWVuEdgSMAAArM/kHwAAIsQ/AABEiH8AAIgQ/wAAEPHjC79zLAEDwLrm/tbyPRqFV2byDwAAEeIfAAAixD8AAESIfwAAiNiPMcaSFy6nw1pn+S0LNgAA8Kvr+Xj3syb/AAAQIf4BACBC/AMAQIT4BwCACPEPAAAR09YHWGLu0+RuAAIAgPuY/AMAQIT4BwCACPEPAAAR4h8AACJeauF3jiVgAAC4j8k/AABEiH8AAIgQ/wAAECH+AQAg4uUXfudYAgaAX839beR7NAWvzOQfAAAixD8AAESIfwAAiBD/AAAQsR9jjCUvXE6Htc6yCUs7AAC8suv5ePezJv8AABAh/gEAIEL8AwBAhPgHAICIt/zC7xK+BgwAQIXJPwAARIh/AACIEP8AABAh/gEAICK/8DvHEjAAAO/I5B8AACLEPwAARIh/AACIEP8AABBh4fdOloABeGVzf8f4Hn//eWUm/wAAECH+AQAgQvwDAECE+AcAgAjxDwAAEfsxxljywuV0WOssb8ENAAAA/KTr+Xj3syb/AAAQIf4BACBC/AMAQIT4BwCAiGnrA7ybuc+nWwIGAOAZmPwDAECE+AcAgAjxDwAAEeIfAAAiLPz+gLkl4N3OIjAAAD/L5B8AACLEPwAARIh/AACIEP8AABBh4XdDvgYMwE+5dfkEy/lbzSsz+QcAgAjxDwAAEeIfAAAixD8AAETsxxhjyQuX02Gts3CDxSIAAG65no93P2vyDwAAEeIfAAAixD8AAESIfwAAiPCF3xfgS8AAADyCyT8AAESIfwAAiBD/AAAQIf4BACDCwu+LsgQMAMBSJv8AABAh/gEAIEL8AwBAhPgHAIAIC79vxBIwALfM/Y3ge/xt5ZWZ/AMAQIT4BwCACPEPAAAR4h8AACLEPwAAROzHGGPJC5fTYa2z8IPcVAAA8B6u5+Pdz5r8AwBAhPgHAIAI8Q8AABHiHwAAIqatD8A25j7zbgkYAOC9mfwDAECE+AcAgAjxDwAAEeIfAAAiLPzyD0vAAADvzeQfAAAixD8AAESIfwAAiBD/AAAQYeGXf2UJGOA9zP2e8z3+DvLKTP4BACBC/AMAQIT4BwCACPEPAAAR+zHGWPLC5XRY6yy8MMtPAADbuJ6Pdz9r8g8AABHiHwAAIsQ/AABEiH8AAIjwhV8ewpeAAQCen8k/AABEiH8AAIgQ/wAAECH+AQAgwsIvq5lbAt7tLAIDAGzF5B8AACLEPwAARIh/AACIEP8AABAh/gEAIMJtP/y4uVuA3AAEsK5bN7CxnL9ZvDKTfwAAiBD/AAAQIf4BACBC/AMAQMR+jDGWvHA5HdY6C/zCQhUAwO9dz8e7nzX5BwCACPEPAAAR4h8AACLEPwAARPjCL0/Ll4ABAB7L5B8AACLEPwAARIh/AACIEP8AABBh4ZeXYgkYAOD7TP4BACBC/AMAQIT4BwCACPEPAAARFn55eZaAAX5v7reS7/E3hldm8g8AABHiHwAAIsQ/AABEiH8AAIjYjzHGkhcup8NaZ4HVWdICAN7N9Xy8+1mTfwAAiBD/AAAQIf4BACBC/AMAQIQv/JLia8AAQJnJPwAARIh/AACIEP8AABAh/gEAIMLCL3mWgAGACpN/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAIt/3ADDcAAe9m7neN7/H3gFdm8g8AABHiHwAAIsQ/AABEiH8AAIjYjzHGkhcup8NaZ4GXY+kLANja9Xy8+1mTfwAAiBD/AAAQIf4BACBC/AMAQIQv/MIf8CVgAOCVmPwDAECE+AcAgAjxDwAAEeIfAAAiLPzCg80tAe92FoEBgO2Z/AMAQIT4BwCACPEPAAAR4h8AACIs/MIP8TVgYEu3LiNgOb/dvDKTfwAAiBD/AAAQIf4BACBC/AMAQMR+jDGWvHA5HdY6C7CzSAYALHM9H+9+1uQfAAAixD8AAESIfwAAiBD/AAAQ4Qu/8GR8CRgAWIvJPwAARIh/AACIEP8AABAh/gEAIEL8AwBAhNt+4AW4AQgAeASTfwAAiBD/AAAQIf4BACBC/AMAQISFX3hRloCBJeZ+M/gev7W8MpN/AACIEP8AABAh/gEAIEL8AwBAxH6MMZa8cDkd1joLsBLLaQDwvq7n493PmvwDAECE+AcAgAjxDwAAEeIfAAAifOEXAnwNGADY7Uz+AQAgQ/wDAECE+AcAgAjxDwAAERZ+IcoSMAD0mPwDAECE+AcAgAjxDwAAEeIfAAAiLPwC/7AEDO9r7t83bMnfl22Y/AMAQIT4BwCACPEPAAAR4h8AACL2Y4yx5IXp43OtswAbu3ch0JIWADyP6/l497Mm/wAAECH+AQAgQvwDAECE+AcAgAhf+AUW8yVgAHhNJv8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABFu+wEeYu4GoN3OLUAA8ExM/gEAIEL8AwBAhPgHAIAI8Q8AABH7McZY8sLldFjrLECEJWD4ebeW8mEr/hY8zvV8vPtZk38AAIgQ/wAAECH+AQAgQvwDAEDE4oXf6eNzrbMAG9tyIdDiFwB8j4VfAADgC/EPAAAR4h8AACLEPwAARExbHwBgt5tfNrYEDACPZfIPAAAR4h8AACLEPwAARIh/AACIsPALPC1LwADwWCb/AAAQIf4BACBC/AMAQIT4BwCAiP0YYyx54XI6rHUWgG+xBAy/N7dAD1vy2/041/Px7mdN/gEAIEL8AwBAhPgHAIAI8Q8AABGLF36nj8+1zgJs7N0WAi2TAVBg4RcAAPhC/AMAQIT4BwCACPEPAAAR4h8AACKmrQ8AsJa524vcAARAmck/AABEiH8AAIgQ/wAAECH+AQAgwsIvkGIJGIAyk38AAIgQ/wAAECH+AQAgQvwDAEDEfowxlrxwOR3WOgvA07AEzLuZW3aHLfmdfZzr+Xj3syb/AAAQIf4BACBC/AMAQIT4BwCAiMULv9PH51pnATZmIfDfWU4D4BlZ+AUAAL4Q/wAAECH+AQAgQvwDAEDEtPUBAF7F3EK0JWAAXonJPwAARIh/AACIEP8AABAh/gEAIMLCL8AfuPVVZIvAADwjk38AAIgQ/wAAECH+AQAgQvwDAEDEfowxlrxwOR3WOgvAW7MEzJZuLafDVvwmPs71fLz7WZN/AACIEP8AABAh/gEAIEL8AwBAxOKF3+njc62zABuzEPjzLLwB8Kcs/AIAAF+IfwAAiBD/AAAQIf4BACBC/AMAQMS09QEAyuZuWHIDEABrMfkHAIAI8Q8AABHiHwAAIsQ/AABEWPgFeDKWgAFYi8k/AABEiH8AAIgQ/wAAECH+AQAgYj/GGEteuJwOa50FgAUsAbPE3CI5bMlv2ONcz8e7nzX5BwCACPEPAAAR4h8AACLEPwAARCxe+J0+Ptc6C7AxC4HvwRIdQIuFXwAA4AvxDwAAEeIfAAAixD8AAERMWx8AgMeaW9y2BAzAbmfyDwAAGeIfAAAixD8AAESIfwAAiLDwCxBgCRiA3c7kHwAAMsQ/AABEiH8AAIgQ/wAAELEfY4wlL1xOh7XOAsDGLAG/r7mlb9iS35vHuZ6Pdz9r8g8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAELH4tp/p43OtswAbcxsIc9zIAfDc3PYDAAB8If4BACBC/AMAQIT4BwCAiGnrAwDw3OYWwS0BA7wmk38AAIgQ/wAAECH+AQAgQvwDAECEhV8AFrv1NWiLwADPzeQfAAAixD8AAESIfwAAiBD/AAAQsR9jjCUvXE6Htc4CwBuyBPwcbi1pw1b8NjzO9Xy8+1mTfwAAiBD/AAAQIf4BACBC/AMAQMTihd/p43OtswAbsxDIT7HoB/A4Fn4BAIAvxD8AAESIfwAAiBD/AAAQMW19AAB65pbLLQEDrM/kHwAAIsQ/AABEiH8AAIgQ/wAAEGHhF4CnYAkYYH0m/wAAECH+AQAgQvwDAECE+AcAgIj9GGMseeFyOqx1FgD4LUvA3zO3UA1b8m/5ca7n493PmvwDAECE+AcAgAjxDwAAEeIfAAAixD8AAEQsvu1n+vhc6yzAxtwGwitzcwhQ5bYfAADgC/EPAAAR4h8AACLEPwAARExbHwAAHmFuYd0SMMCvTP4BACBC/AMAQIT4BwCACPEPAAARFn4BeFuWgAF+ZfIPAAAR4h8AACLEPwAARIh/AACI2I8xxpIXLqfDWmcBgE0UloDnlp9hS4V/dz/lej7e/azJPwAARIh/AACIEP8AABAh/gEAIGLxwu/08bnWWYCNWQiE/2cZEXgVFn4BAIAvxD8AAESIfwAAiBD/AAAQMW19AAB4RnML8JaAgVdn8g8AABHiHwAAIsQ/AABEiH8AAIiw8AsAd7r1FWyLwMCrMPkHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIjYjzHGkhcup8NaZwGAt/FsNwDduqkItvJs/0Ze2fV8vPtZk38AAIgQ/wAAECH+AQAgQvwDAEDE4oXf6eNzrbMAG7MQCOuy4AiswcIvAADwhfgHAIAI8Q8AABHiHwAAIqatDwAAFXNL9ZaAgZ9k8g8AABHiHwAAIsQ/AABEiH8AAIiw8AsAG7IEDPwkk38AAIgQ/wAAECH+AQAgQvwDAEDEfowxlrxwOR3WOgsAcMOfLgHPLRbDliy2P871fLz7WZN/AACIEP8AABAh/gEAIEL8AwBAxOKF3+njc62zABuzEAivx9IkYOEXAAD4QvwDAECE+AcAgAjxDwAAEdPWBwAAvm9uUd8SMHCLyT8AAESIfwAAiBD/AAAQIf4BACDCwi8AvBlLwMAtJv8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABH7McZY8sL08bnWWQCAlbgBCN7X9Xy8+1mTfwAAiBD/AAAQIf4BACBC/AMAQMS09QEAgG1YAoYek38AAIgQ/wAAECH+AQAgQvwDAECEhV8A4B9zS8C7nUVgeBcm/wAAECH+AQAgQvwDAECE+AcAgAgLvwDAb/kaMLwHk38AAIgQ/wAAECH+AQAgQvwDAECEhV8A4FssAcPrMfkHAIAI8Q8AABHiHwAAIsQ/AABEWPgFAB7GEjA8N5N/AACIEP8AABAh/gEAIEL8AwBAhIVfAGBVloDheZj8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEuO0HAPhxbgCCbZj8AwBAhPgHAIAI8Q8AABHiHwAAIiz8AgBPwRIwrM/kHwAAIsQ/AABEiH8AAIgQ/wAAEGHhFwB4WnNLwLudRWD4LpN/AACIEP8AABAh/gEAIEL8AwBAhIVfAODl+BowfI/JPwAARIh/AACIEP8AABAh/gEAIMLCLwDwFiwBw++Z/AMAQIT4BwCACPEPAAAR4h8AACIs/AIAb8sSMPzK5B8AACLEPwAARIh/AACIEP8AABBh4RcASLEETJnJPwAARIh/AACIEP8AABAh/gEAIMLCLwCQZwmYCpN/AACIEP8AABAh/gEAIEL8AwBAhPgHAIAIt/0AAMyYuwFot3MLEK/N5B8AACLEPwAARIh/AACIEP8AABCxeOH31vIL8PossQH83lwL+f3kVZj8AwBAhPgHAIAI8Q8AABHiHwAAIvZjjLHkhenjc62zAAA8vXsvP7EEzE+5no93P2vyDwAAEeIfAAAixD8AAESIfwAAiFj8hV8AAH7Pl4B5Rib/AAAQIf4BACBC/AMAQIT4BwCACAu/AAA/xBIwWzP5BwCACPEPAAAR4h8AACLEPwAAROzHGGPJC5fTYa2zABuzdAbva27RlOfl95glrufj3c+a/AMAQIT4BwCACPEPAAAR4h8AACIWL/xOH59rnQUA4OltuTxtEZg5Fn4BAIAvxD8AAESIfwAAiBD/AAAQIf4BACBi2voAAADcZ+6mITcAsYTJPwAARIh/AACIEP8AABAh/gEAIMLCLwDAC7MEzBIm/wAAECH+AQAgQvwDAECE+AcAgIj9GGMseeFyOqx1FmBjFsTgfc0thdLiN/59Xc/Hu581+QcAgAjxDwAAEeIfAAAixD8AAEQsXvidPj7XOgsAwNN7p+VpS8DvwcIvAADwhfgHAIAI8Q8AABHiHwAAIqatDwAAwDbmlpctAb83k38AAIgQ/wAAECH+AQAgQvwDAECEhV8AAP5x6wvGFoHfg8k/AABEiH8AAIgQ/wAAECH+AQAgYj/GGEteuJwOa50F2JhlLnhft5Y44U/4u/Ecrufj3c+a/AMAQIT4BwCACPEPAAAR4h8AACIWL/xOH59rnQUA4OlZnv53loB/noVfAADgC/EPAAAR4h8AACLEPwAARIh/AACImLY+AAAA72PuNiQ3AD0Pk38AAIgQ/wAAECH+AQAgQvwDAECEhV8AAFZlCfh5mPwDAECE+AcAgAjxDwAAEeIfAAAi9mOMseSFy+mw1lkAAAizBPw91/Px7mdN/gEAIEL8AwBAhPgHAIAI8Q8AABGLF36nj8+1zgIAQMTcV39vsQj87yz8AgAAX4h/AACIEP8AABAh/gEAIGLa+gAAAPBv5paDLQF/j8k/AABEiH8AAIgQ/wAAECH+AQAgwsIvAAAvxxLw95j8AwBAhPgHAIAI8Q8AABHiHwAAIvZjjLHkhcvpsNZZAADgoQpLwNfz8e5nTf4BACBC/AMAQIT4BwCACPEPAAAR4h8AACIW3/YzfXyudRYAACL+8/dfm/2/3+0GILf9AAAAX4h/AACIEP8AABAh/gEAIGLa+gAAAPCT5paN320J+BaTfwAAiBD/AAAQIf4BACBC/AMAQISFXwAA8m59cfjdFoFN/gEAIEL8AwBAhPgHAIAI8Q8AABH7McZY8sLldFjrLAAA8PSebQn4ej7e/azJPwAARIh/AACIEP8AABAh/gEAIGLxwu/08bnWWQAAiLj1Rd1XteUSsIVfAADgC/EPAAAR4h8AACLEPwAARExbHwAAAF7d3ALzs30JeLcz+QcAgAzxDwAAEeIfAAAixD8AAERY+AUAgBU84xKwyT8AAESIfwAAiBD/AAAQIf4BACBiP8YYS164nA5rnQUAAHL+dAn4ej7e/azJPwAARIh/AACIEP8AABAh/gEAIEL8AwBAxOLbfqaPz7XOAgBAxH/+/mvrIzy9e28BctsPAADwhfgHAIAI8Q8AABHiHwAAIqatDwAAAHw1txR97xLwLSb/AAAQIf4BACBC/AMAQIT4BwCACAu/AADwIv70y8gm/wAAECH+AQAgQvwDAECE+AcAgIj9GGNsfQgAAGB9Jv8AABAh/gEAIEL8AwBAhPgHAIAI8Q8AABHiHwAAIsQ/AABEiH8AAIgQ/wAAECH+AQAgQvwDAEDE/wHqMHl1GMmAfAAAAABJRU5ErkJggg==",765      "text/plain": [766       "<Figure size 960x960 with 1 Axes>"767      ]768     },769     "metadata": {},770     "output_type": "display_data"771    }772   ],773   "source": [774    "# mask1 = mask.gather(dim=0, index=result[:, None].expand(mask.shape))\n",775    "# plot_mask(mask1)\n",776    "\n",777    "mask2 = mask.gather(dim=0, index=result[:, None].expand(mask.shape))\n",778    "plot_mask(mask2)\n",779    "\n",780    "\n",781    "mask3 = mask2.gather(dim=1, index=result[None, :].expand(mask.shape))\n",782    "plot_mask(mask3)"783   ]784  },785  {786   "cell_type": "code",787   "execution_count": null,788   "metadata": {},789   "outputs": [],790   "source": [791    "# mask1 = mask.gather(dim=0, index=result[:, None].expand(mask.shape))\n",792    "# plot_mask(mask1)\n",793    "\n",794    "mask2 = mask.gather(dim=1, index=result[None, :].expand(mask.shape))\n",795    "plot_mask(mask2)\n",796    "\n",797    "mask3 = mask2.gather(dim=0, index=result[:, None].expand(mask.shape))\n",798    "plot_mask(mask3)"799   ]800  },801  {802   "cell_type": "code",803   "execution_count": null,804   "metadata": {},805   "outputs": [],806   "source": [807    "# mask1 = mask.gather(dim=0, index=result[None, :].expand(mask.shape))\n",808    "# plot_mask(mask1)\n",809    "\n",810    "mask1 = mask.gather(dim=0, index=result.unsqueeze(1))\n",811    "plot_mask(mask1)\n",812    "\n",813    "# mask2 = mask.gather(dim=1, index=result[:, None].expand(mask.shape))\n",814    "# plot_mask(mask2)"815   ]816  },817  {818   "cell_type": "code",819   "execution_count": null,820   "metadata": {},821   "outputs": [],822   "source": [823    "import torch\n",824    "\n",825    "def get_grid_indices_efficient(q_len, section_size=256, stride=64, gap_size=1, shift=0):\n",826    "    base_indices = torch.arange(0, section_size, stride)\n",827    "    n_full_sections = q_len // (section_size + gap_size)\n",828    "    section_offsets = torch.arange(n_full_sections + 1) * (section_size + gap_size)\n",829    "    indices = (base_indices.view(-1, 1) + section_offsets.view(1, -1) + shift).flatten()\n",830    "    return indices[indices < q_len]\n",831    "\n",832    "get_grid_indices_efficient(256, 65, 8, 0, 0).sort().values"833   ]834  },835  {836   "cell_type": "code",837   "execution_count": null,838   "metadata": {},839   "outputs": [],840   "source": [841    "import json\n",842    "from collections import Counter\n",843    "\n",844    "threshold = -5\n",845    "\n",846    "with open(\"mminference_best_recalls_longvila.json\", \"r\") as f:\n",847    "    data = json.load(f)\n",848    "\n",849    "with open(\"mminference_best_patterns_longvila.json\", \"r\") as f:\n",850    "    bpatterns = json.load(f)\n",851    "\n",852    "gaps = []\n",853    "patterns = []\n",854    "original_patterns = []\n",855    "for layer, heads in data.items():\n",856    "    patterns_p_layers = []\n",857    "\n",858    "    for head, recalls in heads.items():\n",859    "        best_pattern = bpatterns[layer][head]\n",860    "        if best_pattern is None:\n",861    "            continue\n",862    "        best_pattern_type = best_pattern[0]\n",863    "        original_patterns.append(best_pattern_type)\n",864    "        if best_pattern_type != \"grid_attn\":\n",865    "            patterns.append(best_pattern_type)\n",866    "            patterns_p_layers.append(best_pattern_type)\n",867    "\n",868    "            continue\n",869    "        \n",870    "        # Find the specific patterns we want to compare\n",871    "        pattern_182 = next((v for k, v in recalls.items() if k.startswith(\"grid_attn_257_True_True\")), None)\n",872    "        pattern_14 = next((v for k, v in recalls.items() if k.startswith(\"grid_attn_16_True_True\")), None)\n",873    "        \n",874    "        diff = pattern_182 - pattern_14\n",875    "        diff *= 100\n",876    "        gaps.append(diff)\n",877    "        # print(f\"{layer}, Head {head}:\")\n",878    "        # print(f\"Pattern 182: {pattern_182:.4f}\")\n",879    "        # print(f\"Pattern 14: {pattern_14:.4f}\")\n",880    "        # print(f\"Difference (182 - 14): {diff:.4f}\")\n",881    "        # print(\"-\" * 40)\n",882    "\n",883    "        if diff > threshold:\n",884    "            patterns.append(\"grid_attn_257_True_True\")\n",885    "            patterns_p_layers.append(\"grid_attn_257_True_True\")\n",886    "        else:\n",887    "            patterns.append(\"grid_attn_16_True_True\")\n",888    "            patterns_p_layers.append(\"grid_attn_16_True_True\")\n",889    "    \n",890    "    print()\n",891    "    print(layer)\n",892    "    print(Counter(patterns_p_layers))\n",893    "\n",894    "# the ratio of carious gaps\n",895    "# > -1\n",896    "print(f\"ratio of gaps > -1: {len([g for g in gaps if g > -1]) / len(gaps)}\")\n",897    "# > -5\n",898    "print(f\"ratio of gaps > -5: {len([g for g in gaps if g > -5]) / len(gaps)}\")\n",899    "# > -10\n",900    "print(f\"ratio of gaps > -10: {len([g for g in gaps if g > -10]) / len(gaps)}\")\n",901    "# > -20\n",902    "print(f\"ratio of gaps > -20: {len([g for g in gaps if g > -20]) / len(gaps)}\")\n",903    "\n",904    "\n",905    "from collections import Counter\n",906    "# Counter(patterns)\n",907    "print(Counter(original_patterns))"908   ]909  },910  {911   "cell_type": "code",912   "execution_count": null,913   "metadata": {},914   "outputs": [],915   "source": [916    "import json\n",917    "import random\n",918    "import uuid\n",919    "from transformers import AutoTokenizer\n",920    "\n",921    "tok = AutoTokenizer.from_pretrained(\"vision_niah/model_weights/longvila_qwen2_7b_1m/llm\")\n",922    "\n",923    "# generate KVs\n",924    "def generate_a_kv_pair(num_kvs=2500):\n",925    "    kv_pairs = {}\n",926    "    for _ in range(num_kvs):\n",927    "        key = str(uuid.uuid4())\n",928    "        value = str(uuid.uuid4())\n",929    "        kv_pairs[key] = value\n",930    "    return kv_pairs\n",931    "\n",932    "def evenly_select_target_kvs(dic):\n",933    "    keys = list(dic.keys())\n",934    "    total_keys = len(keys)\n",935    "    step = max(1, total_keys // 5)\n",936    "    start = random.randint(0, step - 1)\n",937    "    selected_keys = keys[start::step][:5]\n",938    "    return {key: dic[key] for key in selected_keys}\n",939    "\n",940    "len(tok.encode(json.dumps(generate_a_kv_pair(200))))\n",941    "# dataset = []\n",942    "# for i in range(100):\n",943    "#     kv_pairs = generate_a_kv_pair()\n",944    "#     target_kvs = evenly_select_target_kvs(kv_pairs)\n",945    "    \n",946    "#     context = f\"JSON data:\\n{json.dumps(kv_pairs)}\\n\\n\"\n",947    "#     multi_turns = [\n",948    "#         {\n",949    "#             'input': f\"The key is \\\"{key}\\\". The value associated with the above key is: \",\n",950    "#             'answer': target_kvs[key]\n",951    "#         }\n",952    "#         for key in target_kvs\n",953    "#     ]\n",954    "#     random.shuffle(multi_turns)\n",955    "#     assert len(multi_turns) == 5\n",956    "\n",957    "#     dataset.append({\n",958    "#         'context': context,\n",959    "#         'multi_turns': multi_turns\n",960    "#     })"961   ]962  },963  {964   "cell_type": "code",965   "execution_count": null,966   "metadata": {},967   "outputs": [],968   "source": [969    "kv_pairs[0]"970   ]971  },972  {973   "cell_type": "code",974   "execution_count": null,975   "metadata": {},976   "outputs": [],977   "source": [978    "def calculate_mean_per_layer(s):\n",979    "    layers = {}\n",980    "    current_layer = None\n",981    "    \n",982    "    for line in s.split('\\n'):\n",983    "        stripped = line.strip()\n",984    "        if stripped.startswith('layer_idx:'):\n",985    "            # Extract the layer number\n",986    "            try:\n",987    "                current_layer = int(stripped.split(':')[1].strip())\n",988    "                if current_layer not in layers:\n",989    "                    layers[current_layer] = []\n",990    "            except (IndexError, ValueError):\n",991    "                current_layer = None  # Skip invalid layer lines\n",992    "        elif stripped.startswith('[') and stripped.endswith(']'):\n",993    "            if current_layer is not None:\n",994    "                # Extract the first element\n",995    "                content = stripped[1:-1].strip()\n",996    "                parts = content.split()\n",997    "                if parts:\n",998    "                    try:\n",999    "                        num = float(parts[0])\n",1000    "                        layers[current_layer].append(num)\n",1001    "                    except ValueError:\n",1002    "                        pass  # Skip invalid entries\n",1003    "    \n",1004    "    # Calculate the mean for each layer\n",1005    "    result = {}\n",1006    "    for layer in sorted(layers.keys()):\n",1007    "        elements = layers[layer]\n",1008    "        if elements:\n",1009    "            result[layer] = sum(elements) / len(elements)\n",1010    "        else:\n",1011    "            result[layer] = None  # Handle empty layers\n",1012    "    \n",1013    "    return result\n",1014    "\n",1015    "# Example usage:\n",1016    "s = \"\"\"\n",1017    "\n",1018    "layer_idx: 0\n",1019    "[114.5 25.9]\n",1020    "[8.4 4.4]\n",1021    "[62.5 21.2]\n",1022    "[3.6 7.9]\n",1023    "[113.2 20.1]\n",1024    "[4.8 5.6]\n",1025    "[1.3 4.0]\n",1026    "[10.6 3.6]\n",1027    "[48.2 25.3]\n",1028    "[76.8 23.5]\n",1029    "[21.5 21.1]\n",1030    "[60.6 10.0]\n",1031    "[66.4 16.2]\n",1032    "[84.8 13.7]\n",1033    "[125.5 5.8]\n",1034    "[22.4 27.9]\n",1035    "[7.3 4.4]\n",1036    "[81.3 11.4]\n",1037    "[100.4 19.1]\n",1038    "[60.7 6.8]\n",1039    "[12.7 13.9]\n",1040    "[28.2 10.3]\n",1041    "[68.2 13.3]\n",1042    "[16.2 4.4]\n",1043    "[14.7 4.0]\n",1044    "[21.1 3.7]\n",1045    "[15.9 4.2]\n",1046    "[20.8 3.9]\n",1047    "\n",1048    "layer_idx: 1\n",1049    "[2.7 1.8]\n",1050    "[131.2 7.3]\n",1051    "[10.8 5.2]\n",1052    "[3.4 2.1]\n",1053    "[9.3 3.6]\n",1054    "[4.0 1.5]\n",1055    "[68.0 14.8]\n",1056    "[129.1 9.4]\n",1057    "[67.8 28.1]\n",1058    "[154.8 4.9]\n",1059    "[105.2 5.1]\n",1060    "[118.6 6.0]\n",1061    "[87.7 8.7]\n",1062    "[123.9 5.4]\n",1063    "[73.9 19.2]\n",1064    "[16.5 10.1]\n",1065    "[133.5 10.8]\n",1066    "[16.5 7.4]\n",1067    "[106.3 34.4]\n",1068    "[86.5 17.1]\n",1069    "[43.5 22.9]\n",1070    "[155.8 3.4]\n",1071    "[145.8 4.3]\n",1072    "[53.6 13.6]\n",1073    "[22.3 21.3]\n",1074    "[98.9 23.5]\n",1075    "[147.4 4.5]\n",1076    "[74.9 41.3]\n",1077    "\n",1078    "layer_idx: 2\n",1079    "[60.1 30.7]\n",1080    "[30.6 38.8]\n",1081    "[15.3 26.5]\n",1082    "[50.2 34.1]\n",1083    "[80.3 30.0]\n",1084    "[37.0 36.9]\n",1085    "[4.7 15.4]\n",1086    "[33.0 22.6]\n",1087    "[30.1 25.1]\n",1088    "[35.5 28.5]\n",1089    "[32.4 22.3]\n",1090    "[58.6 28.0]\n",1091    "[56.2 28.9]\n",1092    "[65.0 32.8]\n",1093    "[105.1 17.2]\n",1094    "[35.0 30.9]\n",1095    "[18.0 20.6]\n",1096    "[74.5 29.7]\n",1097    "[77.6 30.4]\n",1098    "[71.9 33.9]\n",1099    "[40.1 42.0]\n",1100    "[110.6 19.6]\n",1101    "[117.8 17.7]\n",1102    "[27.6 22.5]\n",1103    "[87.6 25.4]\n",1104    "[55.2 35.0]\n",1105    "[116.9 25.1]\n",1106    "[41.5 33.2]\n",1107    "\n",1108    "layer_idx: 3\n",1109    "[52.1 25.9]\n",1110    "[21.6 21.6]\n",1111    "[62.3 33.6]\n",1112    "[84.7 26.0]\n",1113    "[46.2 23.0]\n",1114    "[50.1 25.2]\n",1115    "[59.1 31.2]\n",1116    "[17.2 13.0]\n",1117    "[96.5 29.1]\n",1118    "[33.2 29.4]\n",1119    "[40.8 27.2]\n",1120    "[31.1 24.9]\n",1121    "[21.0 12.3]\n",1122    "[40.7 30.0]\n",1123    "[62.6 33.2]\n",1124    "[45.1 24.9]\n",1125    "[93.9 28.8]\n",1126    "[64.0 27.9]\n",1127    "[63.1 25.5]\n",1128    "[17.7 18.6]\n",1129    "[29.5 20.6]\n",1130    "[54.1 25.0]\n",1131    "[38.2 25.7]\n",1132    "[55.3 29.8]\n",1133    "[29.1 28.3]\n",1134    "[21.8 29.0]\n",1135    "[50.8 31.8]\n",1136    "[44.8 31.5]\n",1137    "\n",1138    "layer_idx: 4\n",1139    "[107.3 43.7]\n",1140    "[94.3 54.7]\n",1141    "[135.6 16.8]\n",1142    "[135.2 14.4]\n",1143    "[89.4 54.1]\n",1144    "[77.9 47.4]\n",1145    "[98.4 51.9]\n",1146    "[95.9 25.2]\n",1147    "[7.8 8.3]\n",1148    "[8.9 6.5]\n",1149    "[14.1 9.1]\n",1150    "[7.0 6.0]\n",1151    "[28.9 15.7]\n",1152    "[79.3 31.9]\n",1153    "[126.7 30.4]\n",1154    "[139.0 19.6]\n",1155    "[92.1 55.4]\n",1156    "[143.9 8.1]\n",1157    "[123.8 22.5]\n",1158    "[124.0 33.8]\n",1159    "[109.1 42.2]\n",1160    "[120.7 29.1]\n",1161    "[102.0 37.2]\n",1162    "[129.8 11.7]\n",1163    "[103.3 38.1]\n",1164    "[129.6 21.6]\n",1165    "[143.8 8.9]\n",1166    "[114.5 31.4]\n",1167    "\n",1168    "layer_idx: 5\n",1169    "[40.1 33.9]\n",1170    "[41.6 32.6]\n",1171    "[50.4 44.6]\n",1172    "[52.3 43.2]\n",1173    "[52.3 42.9]\n",1174    "[85.3 46.3]\n",1175    "[107.7 27.5]\n",1176    "[62.4 44.5]\n",1177    "[90.0 46.7]\n",1178    "[91.6 39.9]\n",1179    "[48.6 34.9]\n",1180    "[74.6 57.2]\n",1181    "[31.7 28.9]\n",1182    "[18.8 20.8]\n",1183    "[80.3 40.2]\n",1184    "[116.9 46.3]\n",1185    "[86.1 35.8]\n",1186    "[69.6 39.0]\n",1187    "[40.3 45.5]\n",1188    "[88.1 53.1]\n",1189    "[75.6 43.5]\n",1190    "[91.1 34.0]\n",1191    "[48.4 34.0]\n",1192    "[61.7 31.1]\n",1193    "[63.3 38.3]\n",1194    "[49.1 28.3]\n",1195    "[66.0 27.5]\n",1196    "[60.6 38.9]\n",1197    "\n",1198    "layer_idx: 6\n",1199    "[90.3 29.0]\n",1200    "[68.4 31.5]\n",

Showing the first 1,200 of 6267 lines. Download the file for the rest.