647 lines
221 KiB
Plaintext
647 lines
221 KiB
Plaintext
|
|
{
|
||
|
|
"cells": [
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "2111825f",
|
||
|
|
"metadata": {
|
||
|
|
"id": "2111825f"
|
||
|
|
},
|
||
|
|
"source": [
|
||
|
|
"<h1 align='center' dir='rtl' style='color:yellow'>نام طرح</h1>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "1ab61eab",
|
||
|
|
"metadata": {
|
||
|
|
"id": "1ab61eab"
|
||
|
|
},
|
||
|
|
"source": [
|
||
|
|
"<h2 align='center' dir='rtl'>فراخوانی کتابخانه های مورد نیاز</h2>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 1,
|
||
|
|
"id": "b97c8a8b",
|
||
|
|
"metadata": {
|
||
|
|
"id": "b97c8a8b"
|
||
|
|
},
|
||
|
|
"outputs": [
|
||
|
|
{
|
||
|
|
"name": "stderr",
|
||
|
|
"output_type": "stream",
|
||
|
|
"text": [
|
||
|
|
"2024-11-25 21:52:10.162040: I tensorflow/core/platform/cpu_feature_guard.cc:182] This TensorFlow binary is optimized to use available CPU instructions in performance-critical operations.\n",
|
||
|
|
"To enable the following instructions: SSE4.1 SSE4.2 AVX AVX2 FMA, in other operations, rebuild TensorFlow with the appropriate compiler flags.\n"
|
||
|
|
]
|
||
|
|
}
|
||
|
|
],
|
||
|
|
"source": [
|
||
|
|
"import numpy as np\n",
|
||
|
|
"import matplotlib.pyplot as plt\n",
|
||
|
|
"import tensorflow as tf\n",
|
||
|
|
"from tensorflow import keras\n",
|
||
|
|
"from tensorflow.keras.layers import Dense, Input, Flatten, Conv2D\n",
|
||
|
|
"from tensorflow.keras.models import Model\n",
|
||
|
|
"from tensorflow.keras.datasets import mnist"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "79ef7ea0",
|
||
|
|
"metadata": {
|
||
|
|
"id": "79ef7ea0"
|
||
|
|
},
|
||
|
|
"source": [
|
||
|
|
"<h2 align='center' dir='rtl'> لود کردن دیتاست MNIST</h2>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 2,
|
||
|
|
"id": "511bc7e3",
|
||
|
|
"metadata": {
|
||
|
|
"colab": {
|
||
|
|
"base_uri": "https://localhost:8080/"
|
||
|
|
},
|
||
|
|
"id": "511bc7e3",
|
||
|
|
"outputId": "b627dc6a-3c34-49c0-a89e-698f2e4ead7e"
|
||
|
|
},
|
||
|
|
"outputs": [
|
||
|
|
{
|
||
|
|
"name": "stdout",
|
||
|
|
"output_type": "stream",
|
||
|
|
"text": [
|
||
|
|
"Downloading data from https://storage.googleapis.com/tensorflow/tf-keras-datasets/mnist.npz\n",
|
||
|
|
"\u001b[1m11490434/11490434\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m0s\u001b[0m 0us/step\n"
|
||
|
|
]
|
||
|
|
}
|
||
|
|
],
|
||
|
|
"source": [
|
||
|
|
"train_set, test_set = mnist.load_data()"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "8af0ed71",
|
||
|
|
"metadata": {
|
||
|
|
"id": "8af0ed71"
|
||
|
|
},
|
||
|
|
"source": [
|
||
|
|
"<h2 align='center' dir='rtl'>طراحی لایه های استخراج ویژگی و سنجش نسبت ها</h2>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "bb994741",
|
||
|
|
"metadata": {
|
||
|
|
"id": "bb994741"
|
||
|
|
},
|
||
|
|
"source": [
|
||
|
|
"<h3 align='center' dir='rtl'>لایه استخراج ویژگی</h3>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 3,
|
||
|
|
"id": "8494ca6d",
|
||
|
|
"metadata": {
|
||
|
|
"id": "8494ca6d"
|
||
|
|
},
|
||
|
|
"outputs": [],
|
||
|
|
"source": [
|
||
|
|
"# لایه استخراج ویژگی\n",
|
||
|
|
"kernel_initializer = 'normal'\n",
|
||
|
|
"activation = \"relu\"\n",
|
||
|
|
"\n",
|
||
|
|
"dtype = tf.float32\n",
|
||
|
|
"\n",
|
||
|
|
"\n",
|
||
|
|
"def get_module_1():\n",
|
||
|
|
" shape=(28,28,1)\n",
|
||
|
|
" chs=1\n",
|
||
|
|
" inputs = Input(shape)\n",
|
||
|
|
" layer = Conv2D(filters=32*chs, kernel_size=(3,3), activation=activation, kernel_initializer=kernel_initializer)(inputs)\n",
|
||
|
|
" layer = Conv2D(filters=16*chs, kernel_size=(3,3), strides=(2, 2), activation=activation, kernel_initializer=kernel_initializer)(layer)\n",
|
||
|
|
" layer = Conv2D(filters=8*chs, kernel_size=(3,3), activation=activation, kernel_initializer=kernel_initializer)(layer)\n",
|
||
|
|
" layer = Conv2D(filters=4*chs, kernel_size=(3,3), strides=(2, 2), activation='linear', kernel_initializer=kernel_initializer)(layer)\n",
|
||
|
|
" layer = Conv2D(filters=64*2, kernel_size=(4,4), activation='softplus', kernel_initializer=kernel_initializer)(layer)\n",
|
||
|
|
" layer = Flatten()(layer)\n",
|
||
|
|
"\n",
|
||
|
|
" model = Model(inputs, layer)\n",
|
||
|
|
"\n",
|
||
|
|
" return model"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "ed717db9",
|
||
|
|
"metadata": {
|
||
|
|
"id": "ed717db9"
|
||
|
|
},
|
||
|
|
"source": [
|
||
|
|
"<h3 align='center' dir='rtl'>لایه سنجش نسبت ها</h3>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 7,
|
||
|
|
"id": "f63a27eb",
|
||
|
|
"metadata": {
|
||
|
|
"id": "f63a27eb"
|
||
|
|
},
|
||
|
|
"outputs": [],
|
||
|
|
"source": [
|
||
|
|
"# لایه سنجش نسبت ها\n",
|
||
|
|
"kernel_initializer = 'normal'\n",
|
||
|
|
"activation = \"relu\"\n",
|
||
|
|
"\n",
|
||
|
|
"dtype = tf.float32\n",
|
||
|
|
"\n",
|
||
|
|
"\n",
|
||
|
|
"def get_module_2():\n",
|
||
|
|
" shape=(64*2+10,)\n",
|
||
|
|
" chs=1\n",
|
||
|
|
" inputs = Input(shape)\n",
|
||
|
|
" layer = Flatten()(inputs)\n",
|
||
|
|
" layer = Dense(32, activation='relu', use_bias=True, kernel_initializer=kernel_initializer)(layer)\n",
|
||
|
|
" layer = Dense(16, activation='relu', use_bias=True, kernel_initializer=kernel_initializer)(layer)\n",
|
||
|
|
" log_prob = Dense(10, activation='softmax', use_bias=True, kernel_initializer=kernel_initializer)(layer)\n",
|
||
|
|
"\n",
|
||
|
|
" model = Model(inputs, log_prob)\n",
|
||
|
|
"\n",
|
||
|
|
" return model"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "614adfad",
|
||
|
|
"metadata": {
|
||
|
|
"id": "614adfad"
|
||
|
|
},
|
||
|
|
"source": [
|
||
|
|
"<h2 align='center' dir='rtl'>طراحی و ساختن شبکه عصبی</h2>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 8,
|
||
|
|
"id": "6a14ab0a",
|
||
|
|
"metadata": {
|
||
|
|
"id": "6a14ab0a"
|
||
|
|
},
|
||
|
|
"outputs": [],
|
||
|
|
"source": [
|
||
|
|
"def get_m2_input_vector(f0, f1, lbl1):\n",
|
||
|
|
" # ترکیب ویژگی های کلاس مورد مطالعه با کلاس های دیگر\n",
|
||
|
|
"\n",
|
||
|
|
" inp = f0 * f1\n",
|
||
|
|
" return tf.concat([inp, lbl1], axis=1)\n",
|
||
|
|
"\n",
|
||
|
|
"\n",
|
||
|
|
"class InfluenceModel(keras.Model):\n",
|
||
|
|
" def __init__(self, module_1, module_2):\n",
|
||
|
|
" super(InfluenceModel, self).__init__()\n",
|
||
|
|
" # لایه استخراج ویژگی\n",
|
||
|
|
" self.module_1 = module_1\n",
|
||
|
|
" # لایه تمام متصل (سنجش نسبت ها)\n",
|
||
|
|
" self.module_2 = module_2\n",
|
||
|
|
"\n",
|
||
|
|
" self.loss = keras.losses.SparseCategoricalCrossentropy(from_logits=False)\n",
|
||
|
|
" self.metric = keras.metrics.SparseCategoricalAccuracy()\n",
|
||
|
|
"\n",
|
||
|
|
"\n",
|
||
|
|
" def compile(self, optimizer_1, optimizer_2):\n",
|
||
|
|
" super(InfluenceModel, self).compile()\n",
|
||
|
|
" self.optimizer_1 = optimizer_1\n",
|
||
|
|
" self.optimizer_2 = optimizer_2\n",
|
||
|
|
"\n",
|
||
|
|
" def call_at_zero(self, x):\n",
|
||
|
|
" # جدا کردن تصاویر\n",
|
||
|
|
" imgs = x[0]\n",
|
||
|
|
" # جدا کردن برچسب\n",
|
||
|
|
" lbls = tf.one_hot(x[1], 10)\n",
|
||
|
|
"\n",
|
||
|
|
" # استخراج ویژگی های تصاویر\n",
|
||
|
|
" features = self.module_1(imgs)\n",
|
||
|
|
"\n",
|
||
|
|
" n_feat = features.shape[0]\n",
|
||
|
|
" # جدا سازی ویژگی های هر تصویر\n",
|
||
|
|
" features = tf.split(features, n_feat, axis=0)\n",
|
||
|
|
"\n",
|
||
|
|
" # ویژگی های استحراج شده تصویر مورد مطالعه\n",
|
||
|
|
" f0 = features[0]\n",
|
||
|
|
"\n",
|
||
|
|
" res_mtx = []\n",
|
||
|
|
" res_row = []\n",
|
||
|
|
" for j in range(1, n_feat):\n",
|
||
|
|
" # ترکیب هر نمونه با تصویر مورد مطالعه\n",
|
||
|
|
" f1 = features[j]\n",
|
||
|
|
" lbl1 = lbls[j:j+1,:]\n",
|
||
|
|
" inp = get_m2_input_vector(f0, f1, lbl1)\n",
|
||
|
|
"\n",
|
||
|
|
" res_row += [inp]\n",
|
||
|
|
"\n",
|
||
|
|
" # محاسبه نسبت هر نمونه با تصویر مورد مطالعه\n",
|
||
|
|
" res_row = tf.concat(res_row, axis=0)\n",
|
||
|
|
" res_mtx = self.module_2(res_row)\n",
|
||
|
|
"\n",
|
||
|
|
" # میانگین بر روی نمونه ها\n",
|
||
|
|
" res = tf.reduce_sum(res_mtx, axis=0) / (n_feat-1)\n",
|
||
|
|
" return res, res_mtx\n",
|
||
|
|
"\n",
|
||
|
|
" # Feed - Forward\n",
|
||
|
|
" def call(self, x):\n",
|
||
|
|
"\n",
|
||
|
|
" # جدا کردن تصاویر\n",
|
||
|
|
" imgs = x[0]\n",
|
||
|
|
" # جدا کردن برچسب\n",
|
||
|
|
" lbls = tf.one_hot(x[1], 10)\n",
|
||
|
|
"\n",
|
||
|
|
" # استخراج ویژگی های تصاویر\n",
|
||
|
|
" features = self.module_1(imgs)\n",
|
||
|
|
"\n",
|
||
|
|
" # جدا سازی ویژگی های هر تصویر\n",
|
||
|
|
" n_feat = features.shape[0]\n",
|
||
|
|
" features = tf.split(features, n_feat, axis=0)\n",
|
||
|
|
"\n",
|
||
|
|
"\n",
|
||
|
|
" res_mtx = []\n",
|
||
|
|
" # ساخت جفت های ترکیب شده به ازای هر تصویر برای اموزش مدل\n",
|
||
|
|
" for i in range(n_feat):\n",
|
||
|
|
" res_row = []\n",
|
||
|
|
" for j in range(n_feat):\n",
|
||
|
|
" f0 = features[i]\n",
|
||
|
|
" f1 = features[j]\n",
|
||
|
|
" lbl1 = lbls[j:j+1,:]\n",
|
||
|
|
" inp = get_m2_input_vector(f0, f1, lbl1)\n",
|
||
|
|
" res_row += [inp]\n",
|
||
|
|
" res_row = tf.concat(res_row, axis=0)\n",
|
||
|
|
"\n",
|
||
|
|
" # محاسبه سنجش تاثیر حضور تصاویر دیگر در پیش بینی تصویر مورد مطالعه\n",
|
||
|
|
" res_mtx += [self.module_2(res_row)]\n",
|
||
|
|
"\n",
|
||
|
|
" res_mtx = tf.stack(res_mtx)\n",
|
||
|
|
" #\n",
|
||
|
|
" res_mtx = tf.transpose(res_mtx, [2, 1, 0])\n",
|
||
|
|
" #\n",
|
||
|
|
" res_mtx -= tf.linalg.diag(tf.linalg.diag_part(res_mtx))\n",
|
||
|
|
" #\n",
|
||
|
|
" res_mtx = tf.transpose(res_mtx, [2, 1, 0])\n",
|
||
|
|
" #\n",
|
||
|
|
" res = tf.reduce_sum(res_mtx, axis=1) / (n_feat-1)\n",
|
||
|
|
"\n",
|
||
|
|
" return res, res_mtx\n",
|
||
|
|
"\n",
|
||
|
|
" def train_step(self, x):\n",
|
||
|
|
" # تعریف تابع ضرر\n",
|
||
|
|
" loss = keras.losses.SparseCategoricalCrossentropy(from_logits=True)\n",
|
||
|
|
"\n",
|
||
|
|
" # محاسبه گرادیان و اعمال گرادیان\n",
|
||
|
|
" with tf.GradientTape() as tape_1, tf.GradientTape() as tape_2:\n",
|
||
|
|
" output, _ = self.call(x)\n",
|
||
|
|
" loss_batch = loss(x[1], output)\n",
|
||
|
|
" gradient_1 = tape_1.gradient(loss_batch, self.module_1.trainable_variables)\n",
|
||
|
|
" self.optimizer_1.apply_gradients(zip(gradient_1, self.module_1.trainable_variables))\n",
|
||
|
|
" gradient_2 = tape_2.gradient(loss_batch, self.module_2.trainable_variables)\n",
|
||
|
|
" self.optimizer_2.apply_gradients(zip(gradient_2, self.module_2.trainable_variables))\n",
|
||
|
|
"\n",
|
||
|
|
" # اعمال معیار ارزیابی\n",
|
||
|
|
" metric_batch = self.metric(x[1], output)\n",
|
||
|
|
"\n",
|
||
|
|
" return {\"loss\": loss_batch, \"metric\": metric_batch}\n"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "c1a0891a",
|
||
|
|
"metadata": {
|
||
|
|
"id": "c1a0891a"
|
||
|
|
},
|
||
|
|
"source": [
|
||
|
|
"<h2 align='center' dir='rtl'>اموزش مدل</h2>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "67d4458d",
|
||
|
|
"metadata": {
|
||
|
|
"id": "67d4458d"
|
||
|
|
},
|
||
|
|
"source": [
|
||
|
|
"<h3 align='center' dir='rtl'>اتصال لایه ها</h3>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 9,
|
||
|
|
"id": "f6ed77af",
|
||
|
|
"metadata": {
|
||
|
|
"id": "f6ed77af"
|
||
|
|
},
|
||
|
|
"outputs": [],
|
||
|
|
"source": [
|
||
|
|
"m1 = get_module_1()\n",
|
||
|
|
"m2 = get_module_2()\n",
|
||
|
|
"\n",
|
||
|
|
"model = InfluenceModel(m1, m2)"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "16f4ec08",
|
||
|
|
"metadata": {
|
||
|
|
"id": "16f4ec08"
|
||
|
|
},
|
||
|
|
"source": [
|
||
|
|
"<h3 align='center' dir='rtl'>کامپایل کردن مدل</h3>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 10,
|
||
|
|
"id": "0bbddc92",
|
||
|
|
"metadata": {
|
||
|
|
"id": "0bbddc92"
|
||
|
|
},
|
||
|
|
"outputs": [],
|
||
|
|
"source": [
|
||
|
|
"optimizer_1 = tf.optimizers.Adam(learning_rate=1e-3, beta_1=0.9)\n",
|
||
|
|
"optimizer_2 = tf.optimizers.Adam(learning_rate=1e-3, beta_1=0.9)\n",
|
||
|
|
"\n",
|
||
|
|
"m1.compile(optimizer_1)\n",
|
||
|
|
"m2.compile(optimizer_2)\n",
|
||
|
|
"model.compile(optimizer_1, optimizer_2)\n"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "ac325eb7",
|
||
|
|
"metadata": {
|
||
|
|
"id": "ac325eb7"
|
||
|
|
},
|
||
|
|
"source": [
|
||
|
|
"<h3 align='center' dir='rtl'>اموزش مدل</h3>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 11,
|
||
|
|
"id": "c5e948ab",
|
||
|
|
"metadata": {
|
||
|
|
"colab": {
|
||
|
|
"base_uri": "https://localhost:8080/"
|
||
|
|
},
|
||
|
|
"id": "c5e948ab",
|
||
|
|
"outputId": "137ff429-1ccf-4e2d-dc78-6247b58e893f",
|
||
|
|
"scrolled": true
|
||
|
|
},
|
||
|
|
"outputs": [
|
||
|
|
{
|
||
|
|
"name": "stdout",
|
||
|
|
"output_type": "stream",
|
||
|
|
"text": [
|
||
|
|
"Epoch 1/3\n",
|
||
|
|
"\u001b[1m1875/1875\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m127s\u001b[0m 37ms/step - loss: 1.9111 - metric: 0.6572\n",
|
||
|
|
"Epoch 2/3\n",
|
||
|
|
"\u001b[1m1875/1875\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m81s\u001b[0m 37ms/step - loss: 1.6127 - metric: 0.8748\n",
|
||
|
|
"Epoch 3/3\n",
|
||
|
|
"\u001b[1m1875/1875\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m81s\u001b[0m 37ms/step - loss: 1.5797 - metric: 0.8886\n"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"data": {
|
||
|
|
"text/plain": [
|
||
|
|
"<keras.src.callbacks.history.History at 0x78e4b1533730>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
"execution_count": 11,
|
||
|
|
"metadata": {},
|
||
|
|
"output_type": "execute_result"
|
||
|
|
}
|
||
|
|
],
|
||
|
|
"source": [
|
||
|
|
"model.fit(train_set[0], train_set[1], epochs=3, batch_size=32)\n"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "a42879cd",
|
||
|
|
"metadata": {
|
||
|
|
"id": "a42879cd"
|
||
|
|
},
|
||
|
|
"source": [
|
||
|
|
"<h2 align='center' dir='rtl'>اماده سازی نمونه ها و محاسبه اهمیت هر نمونه نسبت به حضور دیگر نمونه ها</h2>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 12,
|
||
|
|
"id": "2a669978",
|
||
|
|
"metadata": {
|
||
|
|
"id": "2a669978"
|
||
|
|
},
|
||
|
|
"outputs": [],
|
||
|
|
"source": [
|
||
|
|
"def get_samples(sample_ind, batch_size=64):\n",
|
||
|
|
"\n",
|
||
|
|
" # اماده سازی نمونه ها\n",
|
||
|
|
" sample_image = test_set[0][sample_ind:sample_ind+batch_size,...]\n",
|
||
|
|
" sample_label = test_set[1][sample_ind:sample_ind+batch_size,...]\n",
|
||
|
|
" return sample_image, sample_label"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 29,
|
||
|
|
"id": "c910c017",
|
||
|
|
"metadata": {
|
||
|
|
"id": "c910c017"
|
||
|
|
},
|
||
|
|
"outputs": [],
|
||
|
|
"source": [
|
||
|
|
"def get_imporance_function(sample_ind):\n",
|
||
|
|
" loss_fnc = tf.keras.losses.SparseCategoricalCrossentropy(reduction='none', from_logits=False)\n",
|
||
|
|
" img, lbl = get_samples(sample_ind, 1)\n",
|
||
|
|
"\n",
|
||
|
|
" def get_batch_importance(batch_ind, batch_size=64):\n",
|
||
|
|
" sample_images, sample_labels = get_samples(batch_ind, batch_size)\n",
|
||
|
|
" sample_images = np.concatenate([img, sample_images])\n",
|
||
|
|
" sample_labels = np.concatenate([lbl, sample_labels])\n",
|
||
|
|
"\n",
|
||
|
|
" probs_vec, probs_mtx = model.call_at_zero([sample_images, sample_labels])\n",
|
||
|
|
" samples = probs_mtx.numpy()\n",
|
||
|
|
" probs_vec = probs_vec.numpy()\n",
|
||
|
|
"\n",
|
||
|
|
" m = samples.mean(axis=0)\n",
|
||
|
|
"\n",
|
||
|
|
" # حذف نمونه مورد مطالعه و مشاهده تاثیر حضور دیگر نمونه و عدم حضور نمونه بر پیس بینی\n",
|
||
|
|
" sample_output = samples.sum(axis=0, keepdims=True) - samples\n",
|
||
|
|
" sample_output = sample_output/(batch_size-1)\n",
|
||
|
|
" loss_with_sample = loss_fnc(lbl, m)\n",
|
||
|
|
" loss_without_sample = loss_fnc(np.array([lbl]*batch_size), sample_output)\n",
|
||
|
|
" sample_importance = loss_with_sample-loss_without_sample\n",
|
||
|
|
" return sample_images[1:,...], sample_importance, probs_vec, samples\n",
|
||
|
|
" return get_batch_importance, img, lbl"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 30,
|
||
|
|
"id": "a4761918",
|
||
|
|
"metadata": {
|
||
|
|
"id": "a4761918"
|
||
|
|
},
|
||
|
|
"outputs": [],
|
||
|
|
"source": [
|
||
|
|
"batch_ind = 5888\n",
|
||
|
|
"batch_ind = 4498\n",
|
||
|
|
"batch_ind = 1522\n",
|
||
|
|
"batch_ind = 625\n",
|
||
|
|
"batch_ind = 9698\n",
|
||
|
|
"\n",
|
||
|
|
"batch_size = 64\n",
|
||
|
|
"\n",
|
||
|
|
"batch_ind = batch_ind+1\n",
|
||
|
|
"\n",
|
||
|
|
"sample_ind = batch_ind-1\n",
|
||
|
|
"importance_function_batch, img, lbl = get_imporance_function(sample_ind)\n",
|
||
|
|
"sample_images, importance, probs_vec, sample_output = importance_function_batch(batch_ind)"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "52cfc350",
|
||
|
|
"metadata": {
|
||
|
|
"id": "52cfc350"
|
||
|
|
},
|
||
|
|
"source": [
|
||
|
|
"<h2 align='center' dir='rtl'>مشاهده تاثیر دیگر نمونه ها بر پیش بینی مدل</h2>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 31,
|
||
|
|
"id": "7b17b3c2",
|
||
|
|
"metadata": {
|
||
|
|
"colab": {
|
||
|
|
"base_uri": "https://localhost:8080/",
|
||
|
|
"height": 486
|
||
|
|
},
|
||
|
|
"id": "7b17b3c2",
|
||
|
|
"outputId": "de0f988f-0e8d-4572-a28e-e3acecca9d1b"
|
||
|
|
},
|
||
|
|
"outputs": [
|
||
|
|
{
|
||
|
|
"data": {
|
||
|
|
"text/plain": [
|
||
|
|
"<matplotlib.legend.Legend at 0x78e4a624c4f0>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
"execution_count": 31,
|
||
|
|
"metadata": {},
|
||
|
|
"output_type": "execute_result"
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"data": {
|
||
|
|
"image/png": "iVBORw0KGgoAAAANSUhEUgAABFYAAAHDCAYAAAAOU54xAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjguMCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy81sbWrAAAACXBIWXMAAA9hAAAPYQGoP6dpAABVuElEQVR4nO3deXiU9bn/8c8smckeSEImLIGwJICyCoKIPdpKpWpp7fFYDvUIxaU/LVgxbtAK1A2wHhBbFCqi2NN6tK219ahFLRY9KgoFY/VIFjZBIAsEyEYmycz8/kiegUACySSZZ5b367pywTzzfGfuJOhFPny/923x+Xw+AQAAAAAAoMOsZhcAAAAAAAAQrghWAAAAAAAAAkSwAgAAAAAAECCCFQAAAAAAgAARrAAAAAAAAASIYAUAAAAAACBABCsAAAAAAAABIlgBAAAAAAAIEMEKAAAAAABAgAhWAAAAAAAAAkSwApzF9u3b9Z3vfEepqamKj4/XiBEj9Mtf/tLssgAAAAAAIcJudgFAqHrrrbc0bdo0jR07VgsXLlRiYqJ27dqlr776yuzSAAAAAAAhwuLz+XxmFwGEmsrKSuXm5uriiy/WH//4R1mtbO4CAAAAAJyJnxaBVrzwwgsqLS3VI488IqvVqpqaGnm9XrPLAgAAAACEGIIVoBV/+9vflJycrAMHDmjo0KFKTExUcnKybrvtNtXV1ZldHgAAAAAgRBCsAK0oLi5WY2Ojvvvd72rq1Kl6+eWXdeONN2rNmjWaPXu22eUBAAAAAEIEPVaAVgwePFi7d+/WrbfeqtWrV/uv33rrrfr1r3+toqIi5eTkmFghAAAAACAUsGMFaEVcXJwkacaMGS2u/+AHP5Akbd68Oeg1AQAAAABCD8EK0Io+ffpIklwuV4vrGRkZkqSjR48GvSYAAAAAQOghWAFaMW7cOEnSgQMHWlw/ePCgJKlXr15BrwkAAAAAEHoIVoBWfP/735ckrVu3rsX1Z555Rna7XZdddpkJVQEAAAAAQo3d7AKAUDR27FjdeOONevbZZ9XY2KhLL71UmzZt0h/+8ActWLDAf1QIAAAAABDdmAoEtKGhoUFLlizRc889p4MHD2rAgAGaM2eO5s2bZ3ZpAAAAAIAQQbACAAAAAAAQIHqsAAAAAAAABIhgBQAAAAAAIEAEKwAAAAAAAAEiWAEAAAAAAAgQwQoAAAAAAECACFYAAAAAAAACZDe7AAAAgO7k9Xp18OBBJSUlyWKxmF0OAAAIEz6fT1VVVerTp4+s1rb3pbQ7WPmm9bouKQwAEH7e9v7B7BKAgB08eFBZWVlmlwEAAMLU/v371a9fvzafZ8cKAACIaElJSZKa/lKUnJxscjUAACBcVFZWKisry/93ibYQrAAAgIhmHP9JTk4mWAEAAB12rqPENK8FAAAAAAAIEMEKAAAAAABAgAhWAAAAAAAAAkSPFQAAAEkej0cNDQ1mlwF0iZiYGNlsNrPLAICoQLACAACims/nU0lJiY4dO2Z2KUCX6tGjhzIzM8/ZdBEA0DkEKwAAIKoZoUpGRobi4+P5IRRhz+fzqba2VmVlZZKk3r17m1wRAEQ2ghUAABC1PB6PP1RJS0szuxygy8TFxUmSysrKlJGRwbEgAOhGNK8FAABRy+ipEh8fb3IlQNcz/lzTOwgAuhfBCgAAiHoc/0Ek4s81AAQHwQoAAAAAAECACFYAAEDQvPfee5o2bZr69Okji8WiP//5z+dcs2nTJl1wwQVyOp0aMmSI1q9f3+11oqXLLrtM8+bNa/f969evV48ePbqtnrPZtGmTLBZLp6c8tedzzs7O1sqVK/2PT/0zvXfvXlksFuXn53f6fQAAoY3mtQAAIGhqamo0evRo3XjjjfrXf/3Xc96/Z88eXX311br11lv1u9/9Ths3btTNN9+s3r17a+rUqUGoGGjb1q1blZCQ0OpzWVlZOnTokNLT0yU1BT5f//rXdfTo0Rah05/+9CfFxMQEo1yg6/z855LNJi1ceOZzDz0keTxN9wBRgmAFAAAEzZVXXqkrr7yy3fevWbNGAwcO1PLlyyVJw4cP1/vvv6/HH3+cYCXK+Hw+eTwe2e2h89fXXr16tfmczWZTZmbmOV8jNTW1K0sCgsNmkxYtavr9qeHKQw81XX/wQXPqAkzCUSAAABCyNm/erClTprS4NnXqVG3evNmkikLHZZddpttvv13z5s1Tz5495XK5tHbtWtXU1Gj27NlKSkrSkCFD9Ne//rXFunfffVcTJkyQ0+lU7969NX/+fDU2Nvqfr6mp0cyZM5WYmKjevXv7Q61Tud1u3X333erbt68SEhI0ceJEbdq0qd21G8dkXnzxRV188cWKjY3ViBEj9O677/rvMY70/PWvf9W4cePkdDr1/vvvy+126yc/+YkyMjIUGxurSy65RFu3bj3jPT744AONGjVKsbGxuuiii/T555/7nzty5IhmzJihvn37Kj4+XiNHjtR///d/n/EajY2Nmjt3rlJSUpSenq6FCxfK5/P5nz/9KFBrn2N+fr727t2rr3/965Kknj17ymKx6Ic//KGkM48Cnetr++WXX2ratGnq2bOnEhISdP755+uNN95oz5cd6DoLFzaFJ4sWNYUpUstQpbWdLEAEI1gBAAAhq6SkRC6Xq8U1l8ulyspKnThxotU1brdblZWVLT46wufzqba+0ZSPU39ob4/nn39e6enp2rJli26//Xbddtttuu6663TxxRdr+/btuuKKK3TDDTeotrZWknTgwAFdddVVuvDCC/Xpp59q9erVWrdunR5++GH/a95zzz1699139Ze//EVvvfWWNm3apO3bt7d437lz52rz5s168cUX9c9//lPXXXedvvWtb6m4uLhD9d9zzz2666679Mknn2jSpEmaNm2ajhw50uKe+fPna9myZdqxY4dGjRqle++9Vy+//LKef/55bd++XUOGDNHUqVNVUVFxxmsvX75cW7duVa9evTRt2jT/2OG6ujqNGzdOr7/+uj7//HP96Ec/0g033KAtW7ac8fW12+3asmWLnnjiCa1YsULPPPNMhz5HqelY0MsvvyxJKiws1KFDh/TEE0+0eu+5vrZz5syR2+3We++9p88++0yPPvqoEhMTO1wT0GmnhitOJ6EKolro7KUEAADoAkuXLtUDDzwQ8PoTDR6dt+jNLqyo/b54cKriHe3/69no0aN1//33S5IWLFigZcuWKT09XbfccoskadGiRVq9erX++c9/6qKLLtJTTz2lrKwsrVq1ShaLRcOGDdPBgwd13333adGiRaqtrdW6dev029/+VpdffrmkpnChX79+/vfct2+fnnvuOe3bt099+vSRJN19993asGGDnnvuOS1ZsqTd9c+dO1fXXnutJGn16tXasGGD1q1bp3vvvdd/z4MPPqhvfvObkpp206xevVrr16/3Hylbu3at3n77ba1bt0733HOPf93ixYv964zP4ZVXXtH3v/999e3bV3fffbf/3ttvv11vvvmmfv/732vChAn+61lZWXr88cdlsVg0dOhQffbZZ3r88cf9X9/2stls/iM/GRkZbTb2bc/Xdt++fbr22ms1cuRISdKgQYM6VAvQpRYulB5+WKqvlxwOQhVELYIVAAAQsjIzM1VaWtriWmlpqZKTkxUXF9fqmgULFigvL8//uLKyUllZWd1ap1lGjRrl/73NZlNaWpr/B25J/t0+ZWVlkqQdO3Zo0qRJslgs/nsmT56s6upqffXVVzp69Kjq6+s1ceJE//OpqakaOnSo//Fnn30mj8ej3NzcFrW43W6lpaV1qP5Jkyb5f2+32zV+/Hjt2LGjxT3jx4/3/37Xrl1qaGjQ5MmT/ddiYmI0YcKEM9ad+trG52Dc4/F4tGTJEv3+97/XgQMHVF9fL7fbrfj4+BavcdFFF7X4Wk2aNEnLly+Xx+ORzWbr0OfaHu352v7kJz/RbbfdprfeektTpkzRtdde2+LPARBUDz10MlSpr296TLiCKESwAgAAQtakSZPO6B/x9ttvt/ih+XROp1NOpzPg94yLsemLB81pjBsX07Ef1k+fJ
|
||
|
|
"text/plain": [
|
||
|
|
"<Figure size 1500x500 with 2 Axes>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
"metadata": {},
|
||
|
|
"output_type": "display_data"
|
||
|
|
}
|
||
|
|
],
|
||
|
|
"source": [
|
||
|
|
"plt.figure(figsize=(15,5))\n",
|
||
|
|
"plt.subplot(121)\n",
|
||
|
|
"plt.imshow(img[0])\n",
|
||
|
|
"plt.title(lbl[0])\n",
|
||
|
|
"plt.axis('off')\n",
|
||
|
|
"\n",
|
||
|
|
"plt.subplot(122)\n",
|
||
|
|
"plt.plot(probs_vec, label='model probabilities')\n",
|
||
|
|
"plt.plot(lbl, 1, 'rx', label='correct label')\n",
|
||
|
|
"plt.legend()"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 32,
|
||
|
|
"id": "e1124039",
|
||
|
|
"metadata": {
|
||
|
|
"colab": {
|
||
|
|
"base_uri": "https://localhost:8080/",
|
||
|
|
"height": 1000
|
||
|
|
},
|
||
|
|
"id": "e1124039",
|
||
|
|
"outputId": "a9579fb9-09e7-467a-f897-73eec99eb1ec"
|
||
|
|
},
|
||
|
|
"outputs": [
|
||
|
|
{
|
||
|
|
"data": {
|
||
|
|
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAiMAAAGdCAYAAADAAnMpAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjguMCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy81sbWrAAAACXBIWXMAAA9hAAAPYQGoP6dpAAAyLElEQVR4nO3dfXBU53n38d/ZlXZXEpIQCMSbsCSnrU1sA4bAyDRNMlVNHZeOO32hjhMYkngmKbTYmqYBx4amjpGdqSmdGJvYMU1nGj8mTRs3rV0yrlLiusaDDaZjP/XLE0tYMkRCL6AVetmVds/zh3RWEpZAK+3ufXbP9zOjGXvZ1V5YBv1039d93ZZt27YAAAAM8ZkuAAAAeBthBAAAGEUYAQAARhFGAACAUYQRAABgFGEEAAAYRRgBAABGEUYAAIBReaYLmI54PK5z586puLhYlmWZLgcAAEyDbdvq7e3VkiVL5PNNvf6RFWHk3LlzqqysNF0GAACYgdbWVi1btmzKX8+KMFJcXCxp5DdTUlJiuBoAADAd4XBYlZWVie/jU8mKMOJszZSUlBBGAADIMldrsaCBFQAAGEUYAQAARhFGAACAUYQRAABgFGEEAAAYRRgBAABGEUYAAIBRhBEAAGAUYQQAABiVdBh56aWXtGnTJi1ZskSWZem555676muOHTumm2++WcFgUB/72Mf0/e9/fwalAgCAXJR0GOnr69PKlSt18ODBaT2/ublZt99+uz7zmc/o9OnTuueee/TlL39ZP/3pT5MuFgAA5J6k76a57bbbdNttt037+YcOHVJ1dbUeffRRSdL111+vl19+WX/zN3+jjRs3Jvv2AAAgx6S9Z+T48eOqq6ub8NjGjRt1/PjxKV8TiUQUDocnfACA1/2fEy165RedpssAUi7tYaStrU0VFRUTHquoqFA4HNbAwMCkr2loaFBpaWnio7KyMt1lAoCr/b/2Xu3+5zf1Z8+eNl0KkHKuPE2ze/du9fT0JD5aW1tNlwQARr3XfkmS1Hkpogt9UcPVAKmVdM9IshYtWqT29vYJj7W3t6ukpEQFBQWTviYYDCoYDKa7NADIGk0dl8b+ufOS1hTNM1gNkFppXxmpra1VY2PjhMdefPFF1dbWpvutASBnNHf2Jf65qaPvCs8Esk/SYeTSpUs6ffq0Tp8+LWnk6O7p06fV0tIiaWSLZcuWLYnnf+UrX1FTU5P+4i/+Qu+8844ef/xx/fCHP9S9996bmt8BAHhA07gwMj6YALkg6TDy+uuva/Xq1Vq9erUkqb6+XqtXr9aePXskSb/85S8TwUSSqqur9fzzz+vFF1/UypUr9eijj+p73/sex3oBYJps256wTUMYQa5Jumfk05/+tGzbnvLXJ5uu+ulPf1pvvPFGsm8FAJB0oX9I4cHhxL8TRpBrXHmaBgAwxlkVCeSN/JXd3NmneHzqHwqBbEMYAQCXc/pF1iwvU77fUmQ4rnM9k89pArIRYQQAXM7ZlvnYwjm6Zn7RhMeAXEAYAQCXax49yltdXqTqcsIIck/ah54BAGbHCR41C4rU3jsoiVkjyC2EEQBwsXjcVnPXaBgpn6P28EgYYWUEuYRtGgBwsbMXBxQdjivfb2lpWYGqy+dIGhkJD+QKwggAuJizAnLN/CL5fVaiZ+TDCwOKDMdMlgakDGEEAFzMCSNOCCmfE1BxME+2LbV09ZssDUgZwggAuFiieXU0jFiWpeoFI//cRN8IcgRhBABcrGncSRpHDcd7kWMIIwDgYs2jjapO4+r4f27meC9yBGEEAFxqcCimDy+MjH13ekYkjdum4UQNcgNhBABcqqW7X7YtFQfzVD4nkHicbRrkGsIIALiUM2W1ekGRLMtKPF41GkY6L0XVMzBkpDYglQgjAOBSl5+kccwJ5mlhcVCSdIbVEeQAwggAuNRkzasO53QNWzXIBYQRAHCpxMCzBUUf+bWxsfCEEWQ/wggAuJTTM3L5Ns34x5o6OFGD7EcYAQAX6ukfUldfVNJYw+p41ZyoQQ4hjACACzV3jYSMhcVBzQnmfeTXq8f1jNi2ndHagFQjjACACznNqzWT9ItI0vJ5hfL7LPVHYzrfG8lkaUDKEUYAwIWcUe+TnaSRpHy/T8vnFUoa6y0BshVhBABcqGmKGSPj0TeCXEEYAQAXSkxfnUYY4UQNsh1hBABcxrbtK84YcbAyglxBGAEAl2kPRzQwFJPfZ6myrHDK53FhHnIFYQQAXKZp9CTN8nmFCuRN/dd0zYKR5taW7n4NxeIZqQ1IB8IIALhMYovmCv0iklRRElRBvl/DcVsfXhjIRGlAWhBGAMBlptO8KkmWZY3rG6GJFdmLMAIALjPdlRFprMGVWSPIZoQRAHCZ5mnMGHEkLsyjiRVZjDACAC4yFIurpbtf0pWP9ToS2zSsjCCLEUYAwEVau/sVi9sqyPdrUUnoqs93TtRwvBfZjDACAC4yvl/EsqyrPr96/sjKSFt4UH2R4bTWBqQLYQQAXCRxkmYaWzSSVFqYr/lFAUmsjiB7EUYAwEWmc0He5RgLj2xHGAEAF3HmhUznWK+DMIJsRxgBABdJZsaIw9nSIYwgWxFGAMAl+iLDag9HJEk15XOm/TrnucwaQbYijACASzgrG/OLAiotzJ/262qclZGOS7JtOy21AelEGAEAl2iawRaNNHK7r2VJ4cFhdfVF01EakFaEEQBwieZpXpB3uVC+X0vnFox8DrZqkIUIIwDgEomTNNOcMTIeY+GRzQgjAOASYxfkTb951cGFechmhBEAcAHbtscGns1gZWTsjppLKa0LyATCCAC4QFdfVL2Dw7KskYbUZDH4DNmMMAIALuDcSbN0boFC+f6kX++EkTNdI7f+AtmEMAIALjCTMfDjLZlboECeT9HhuM5dHEhlaUDaEUYAwAVmckHeeH6fpar5hRM+F5AtCCMA4ALOkVynEXUmxo730sSK7EIYAQAXmMkFeZcbO1HDygiyC2EEAAyLxW190NUvaXZhpJpZI8hShBEAMOzshQFFY3EF8nxaMjrWfSYSg8+YwoosQxgBAMOaRk/SVM0vlN9nzfjzOCsj53oGNDgUS0ltQCYQRgDAsFT0i0jSvKKASkJ5sm0ltn2AbDCjMHLw4EFVVVUpFApp/fr1OnHixBWff+DAAf3ar/2aCgoKVFlZqXvvvVeDg4MzKhgAck3iTppZnKSRJMuyGAuPrJR0GDly5Ijq6+u1d+9enTp1SitXrtTGjRt1/vz5SZ//zDPPaNeuXdq7d6/efvttPf300zpy5Ijuu+++WRcPALkgVSsjEhfmITslHUb279+vu+++W9u2bdOKFSt06NAhFRYW6vDhw5M+/5VXXtGGDRv0uc99TlVVVbr11lt15513XnU1BQC8wmk4nenAs/HGZo0QRpA9kgoj0WhUJ0+eVF1d3dgn8PlUV1en48ePT/qaW265RSdPnkyEj6amJr3wwgv67Gc/O4uyASA3DA7FdHZ0fHsqVkaqF7AyguyTl8yTOzs7FYvFVFFRMeHxiooKvfPOO5O+5nOf+5w6Ozv167/+67JtW8PDw/rKV75yxW2aSCSiSCSS+PdwOJxMmQCQNc50jYSGklCe5hUFZv35uL0X2Sjtp2mOHTumffv26fHHH9epU6f0z//8z3r++ef14IMPTvmahoYGlZaWJj4qKyvTXSYAGOFsp1QvmCPLmvmxXocTRrr7orrYH5315wMyIakwUl5eLr/fr/b29gmPt7e3a9GiRZO+5oEHHtAXvvAFffnLX9aNN96o3/u939O+ffvU0NCgeDw+6Wt2796tnp6exEdra2syZQJA1nC2U65NwRaNJBUG8rS4NCSJ1RFkj6TCSCAQ0Jo1a9TY2Jh4LB6Pq7GxUbW1tZO+pr+/Xz7fxLfx+/2SJNu2J31NMBhUSUnJhA8AyEWpPEnjYKsG2SapnhFJqq+v19atW7V27VqtW7dOBw4cUF9fn7Zt2yZJ2rJli5YuXaqGhgZJ0qZNm7R//36tXr1a69ev1y9+8Qs98MAD2rRpUyKUAIBXNY3esOs0nqZCdXmRXnm/i7HwyBpJh5HNm
|
||
|
|
"text/plain": [
|
||
|
|
"<Figure size 640x480 with 1 Axes>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
"metadata": {},
|
||
|
|
"output_type": "display_data"
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"data": {
|
||
|
|
"image/png": "iVBORw0KGgoAAAANSUhEUgAABI8AAAQqCAYAAADOJF9fAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjguMCwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy81sbWrAAAACXBIWXMAAA9hAAAPYQGoP6dpAAEAAElEQVR4nOzdd1hTZ/sH8DuEvWQvEZSlOHGAq+5VV62KWzvcs9bWUVutVat11YGiYrWuuq3Wau1wvda69ypucQIqiojs5Pn94Y/AnXACYVSSfD/X1evN95wn55zA/Z4cHnPuyIQQggAAAAAAAAAAAPJg8rYPAAAAAAAAAAAASi9MHgEAAAAAAAAAgCRMHgEAAAAAAAAAgCRMHgEAAAAAAAAAgCRMHgEAAAAAAAAAgCRMHgEAAAAAAAAAgCRMHgEAAAAAAAAAgCRMHgEAAAAAAAAAgCRMHgEAAAAAAAAAgCRMHgEAAAAAAAAAgKT/ZPIoMTGRBg8eTK6urmRjY0PNmjWjc+fOFfj50dHR9O6775KtrS05OTlRv3796OnTpxrjlEolzZkzhypUqECWlpZUvXp12rRpk8a4U6dO0fDhw6l27dpkZmZGMplM6/7j4+NpyJAhVLZsWbK0tKTy5cvTgAEDNMY9evSIunfvTg4ODmRvb0+dOnWiO3fu5LnNVatWUXBwMFlaWlJgYCAtXrxYcv9btmyh+vXrk42NDTk4OFCDBg3o4MGDWo/ZkOlzPclksjz/mzVrlsbYzZs3U61atcjS0pJcXV1pwIAB9OzZM41x8fHx9PHHH5ObmxtZWVlRrVq1aNu2bXnuv6DbNCaoJw71VHSoKa6gNbVz505q06YNeXl5kYWFBXl7e1N4eDhduXIlvx+ZQUM95VizZo3kNmUyGW3YsEHnbRoj1FQO1FTRoZ5y6FJPO3bsoB49epCfnx9ZW1tTxYoV6fPPP6fExMQC/+wMEepJU0HnIvbv30/NmjUjFxcXcnBwoLCwMFq/fr22H5fuRAlTKBSiQYMGwsbGRnzzzTdiyZIlonLlysLOzk7cuHEj3+c/ePBAuLi4CH9/f7Fo0SIxY8YM4ejoKGrUqCHS09PZ2C+++EIQkRg0aJBYsWKFaN++vSAisWnTJjZuypQpwszMTNSuXVsEBQUJbT+G+/fvi3Llyoly5cqJadOmiVWrVonp06eLjh07snGvXr0SgYGBws3NTcyePVvMnz9flCtXTnh7e4tnz56xscuXLxdEJLp27SpWrFgh+vXrJ4hIzJo1S2P/U6ZMETKZTHTr1k0sX75cLF68WAwZMkSsW7cu35+dIdL3eiIi0apVK7F+/Xr235UrV9i4pUuXCiISLVq0EJGRkWLixInC2tpaVK9eXaSmpqrGvXz5UgQEBAg7OzsxadIksWTJEtG4cWNBRGLDhg2F2qYxQT2hnoobaqrwNTV16lTRo0cPMWvWLLFy5Urx7bffCj8/P2FlZSUuXLiQ78/OEKGeeD3dvn1bY1vr168XtWrVEnK5XMTGxuq8TWODmkJNFSfUU+HrydnZWVSrVk1MnjxZ/PDDD+KTTz4R5ubmolKlSiIlJSXfn50hQj1pnksKOhexa9cuIZPJRIMGDcTixYvZ9db8+fPz/dkVVIlPHm3ZskUQkdi2bZtq2ZMnT4SDg4Po1atXvs8fNmyYsLKyEvfu3VMt27dvnyAiERUVpVr28OFDYWZmJkaMGKFaplQqRaNGjYS3t7fIyspSLY+Li1P9n3LEiBFai6Bt27aiQoUKGhNA6mbPni2ISJw6dUq1LDo6WsjlcjFx4kTVspSUFOHs7Czat2/Pnt+nTx9hY2Mjnj9/rlp2/PhxIZPJivUXru/0vZ6IiG0zL+np6cLBwUE0btxYKJVK1fLdu3cLIhIRERGqZXPmzBFEJA4cOKBaplAoRGhoqPDw8FCdKHXZpjFBPaGeihtqqnA1JSUuLk6YmpqKIUOGaB1nqFBP+Z9PUlJShJ2dnWjVqlWxbdOQoaZQU8UJ9VS4ehJCiEOHDmmMXbt2rSAi8cMPP2jdpqFCPWnWU0HnIlq1aiW8vLxEWlqaallmZqbw9/cX1atX1/pcXZT45FG3bt2Eu7u7UCgUbPngwYOFtbU1e4F5cXNzE926ddNYHhQUJFq0aKHKkZGRgojE1atX2biNGzcKIhJHjhzJc/vaiiA6OloQkVi6dKkQQojU1FSRkZGR59jQ0FARGhqqsbx169bC399flX/77TdBROK3335j444dOyaISKxfv161rEePHsLT01MoFAqhVCrFq1ev8ty3MdHnehIi56SSkpIi+a9UZ8+eFUQkIiMjNdbZ2tqKBg0aqHLHjh2Fq6urxri5c+cKIhJ//fWXzts0Jqgn1FNxQ00VrqakKJVKYW9vL3r06KF1nKFCPeV/Psn+Y2PNmjXFtk1DhppCTRUn1FPh6klKUlKSICLx2Wef5TvWEKGeeD3pMhdRt25dUaVKlTyX161bV/KYdVXiPY/Onz9PtWrVIhMTvquwsDBKSUmhGzduSD730aNH9OTJE6pTp47GurCwMDp//jzbj42NDQUHB2uMy16vq/379xMRkbu7O7Vo0YKsrKzIysqK2rZtSzExMapxSqWSLl26JHmct2/fplevXrHjUB9bu3ZtMjExYcd54MABCg0NpYiICHJ1dSU7Ozvy9PSkJUuW6PxaDIU+11O2NWvWkI2NDVlZWVHlypVp48aNbH16ejoREVlZWWk818rKis6fP09KpVI1Nq9x1tbWRER09uxZnbdpTFBPqKfihpoqXE3llpiYSE+fPqXLly/TwIEDKSkpiVq0aFHo16PPUE/5n082bNhAVlZW1KVLl2LbpiFDTaGmihPqqXD1JCUuLo6IiFxcXHR5CQYD9cTrqaBzEURETZs2patXr9LkyZPp1q1bdPv2bZo+fTqdOXOGxo8fX+jXo67EJ49iY2PJ09NTY3n2ssePH2t9bu6x6s9//vy56hcQGxtL7u7uGk2sCrIfKTdv3iQiosGDB5O5uTlt2bKFZs2aRf/88w+1bNmSUlJSiIhUx1GQ1xkbG0tyuZzc3NzYOHNzc3J2dlaNe/HiBT179oyOHj1KkydPpi+++IK2bNlCISEhNGrUKIqKitL59RgCfa4nIqIGDRrQjBkz6JdffqFly5aRXC6nPn360LJly1RjAgMDSSaT0dGjR9lzr1+/Tk+fPqXU1FR68eIFERFVrFiRHj58SPfu3WNjjxw5QkRvTqS6btOYoJ5QT8UNNVW4msqtXr165ObmRtWrV6etW7fSpEmT8mwMaQxQT9rPJ8+fP6c//viDOnbsSHZ2dsWyTUOHmkJNFSfUU+HqScrs2bNJLpdTeHh4oV6PvkM98Xoq6FwEEdHkyZOpe/fuNGPGDAoMDKSAgACaNWsW/fzzzwWauCwo02LbkoTU1FSysLDQWG5paalar+25RJTv8y0sLIq0HynJyclEROTh4UG//fabahbU29ubevXqRRs3bqSBAwcW+Diz/9fc3DzP/VlaWqrGZe87ISGBNm/eTD169CAiovDwcKpWrRp9++23NGTIEJ1fk77T53oiIo0TRf/+/al27dr05Zdf0kcffURWVlbk4uJC3bt3p7Vr11JwcDB17tyZHj16RKNGjSIzMzPKzMxU7X/gwIG0fPly6t69Oy1YsIDc3d1p69attHPnTnacumzTmKCeUE/FDTVVuJrKbfXq1ZSUlER37tyh1atXU2pqKikUCo1/iTQGqCft55Pt27dTRkYG9enThy3HOUoaago1VZxQT4Wrp7xs3LiRVq1aRePHj6fAwMBCvR59h3ri9VTQuYjs1x0UFETh4eHUpUsXUigUtGLFCurbty/t27eP6tWrV6jXpK7YrsQyMjIoLi6O/adQKMjKyko1y5dbWloaEeX9ka1s2esK8vyi7Ce//Xfv3p1dtHbr1o1MTU3p2LFjhTrOjIyMPPeXlpbGxhERmZmZsdlnExMT6tGjBz18+JDu37+v82vSF4ZYT3kxNzenkSNHUmJiIrt9Iyoqitq1a0djx44lf39/aty4MVWrVo06duxIRES2trZERFS9e
|
||
|
|
"text/plain": [
|
||
|
|
"<Figure size 1500x1500 with 64 Axes>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
"metadata": {},
|
||
|
|
"output_type": "display_data"
|
||
|
|
}
|
||
|
|
],
|
||
|
|
"source": [
|
||
|
|
"from matplotlib import pyplot as plt\n",
|
||
|
|
"import numpy as np\n",
|
||
|
|
"\n",
|
||
|
|
"inds = np.argsort(importance)\n",
|
||
|
|
"plt.figure()\n",
|
||
|
|
"plt.plot(sample_output[inds[:5]].mean(axis=0))\n",
|
||
|
|
"\n",
|
||
|
|
"\n",
|
||
|
|
"plt.figure(figsize=(15,15))\n",
|
||
|
|
"for i in range(len(inds)):\n",
|
||
|
|
" plt.subplot(batch_size//8+1, 8, i+1)\n",
|
||
|
|
" plt.imshow(sample_images[inds[i],...])\n",
|
||
|
|
" plt.title(\"%.7f\"%(importance[inds[i]]))\n",
|
||
|
|
" plt.axis('off')"
|
||
|
|
]
|
||
|
|
}
|
||
|
|
],
|
||
|
|
"metadata": {
|
||
|
|
"colab": {
|
||
|
|
"provenance": []
|
||
|
|
},
|
||
|
|
"kernelspec": {
|
||
|
|
"display_name": "CV",
|
||
|
|
"language": "python",
|
||
|
|
"name": "python3"
|
||
|
|
},
|
||
|
|
"language_info": {
|
||
|
|
"codemirror_mode": {
|
||
|
|
"name": "ipython",
|
||
|
|
"version": 3
|
||
|
|
},
|
||
|
|
"file_extension": ".py",
|
||
|
|
"mimetype": "text/x-python",
|
||
|
|
"name": "python",
|
||
|
|
"nbconvert_exporter": "python",
|
||
|
|
"pygments_lexer": "ipython3",
|
||
|
|
"version": "3.11.5"
|
||
|
|
}
|
||
|
|
},
|
||
|
|
"nbformat": 4,
|
||
|
|
"nbformat_minor": 5
|
||
|
|
}
|