563 lines
412 KiB
Plaintext
563 lines
412 KiB
Plaintext
|
|
{
|
||
|
|
"cells": [
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "2111825f",
|
||
|
|
"metadata": {},
|
||
|
|
"source": [
|
||
|
|
"<h1 align='center' dir='rtl' style='color:yellow'>نام طرح</h1>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "1ab61eab",
|
||
|
|
"metadata": {},
|
||
|
|
"source": [
|
||
|
|
"<h2 align='center' dir='rtl'>فراخوانی کتابخانه های مورد نیاز</h2>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 11,
|
||
|
|
"id": "b97c8a8b",
|
||
|
|
"metadata": {},
|
||
|
|
"outputs": [],
|
||
|
|
"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": {},
|
||
|
|
"source": [
|
||
|
|
"<h2 align='center' dir='rtl'> لود کردن دیتاست MNIST</h2>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 12,
|
||
|
|
"id": "511bc7e3",
|
||
|
|
"metadata": {},
|
||
|
|
"outputs": [],
|
||
|
|
"source": [
|
||
|
|
"train_set, test_set = mnist.load_data()"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "8af0ed71",
|
||
|
|
"metadata": {},
|
||
|
|
"source": [
|
||
|
|
"<h2 align='center' dir='rtl'>طراحی لایه های استخراج ویژگی و سنجش نسبت ها</h2>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "bb994741",
|
||
|
|
"metadata": {},
|
||
|
|
"source": [
|
||
|
|
"<h3 align='center' dir='rtl'>لایه استخراج ویژگی</h3>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 13,
|
||
|
|
"id": "8494ca6d",
|
||
|
|
"metadata": {},
|
||
|
|
"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": {},
|
||
|
|
"source": [
|
||
|
|
"<h3 align='center' dir='rtl'>لایه سنجش نسبت ها</h3>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 14,
|
||
|
|
"id": "f63a27eb",
|
||
|
|
"metadata": {},
|
||
|
|
"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": {},
|
||
|
|
"source": [
|
||
|
|
"<h2 align='center' dir='rtl'>طراحی و ساختن شبکه عصبی</h2>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 106,
|
||
|
|
"id": "6a14ab0a",
|
||
|
|
"metadata": {},
|
||
|
|
"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": {},
|
||
|
|
"source": [
|
||
|
|
"<h2 align='center' dir='rtl'>اموزش مدل</h2>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "67d4458d",
|
||
|
|
"metadata": {},
|
||
|
|
"source": [
|
||
|
|
"<h3 align='center' dir='rtl'>اتصال لایه ها</h3>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 107,
|
||
|
|
"id": "f6ed77af",
|
||
|
|
"metadata": {},
|
||
|
|
"outputs": [],
|
||
|
|
"source": [
|
||
|
|
"m1 = get_module_1()\n",
|
||
|
|
"m2 = get_module_2()\n",
|
||
|
|
"\n",
|
||
|
|
"model = InfluenceModel(m1, m2)"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "markdown",
|
||
|
|
"id": "16f4ec08",
|
||
|
|
"metadata": {},
|
||
|
|
"source": [
|
||
|
|
"<h3 align='center' dir='rtl'>کامپایل کردن مدل</h3>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 108,
|
||
|
|
"id": "0bbddc92",
|
||
|
|
"metadata": {},
|
||
|
|
"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": {},
|
||
|
|
"source": [
|
||
|
|
"<h3 align='center' dir='rtl'>اموزش مدل</h3>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 109,
|
||
|
|
"id": "c5e948ab",
|
||
|
|
"metadata": {
|
||
|
|
"scrolled": true
|
||
|
|
},
|
||
|
|
"outputs": [
|
||
|
|
{
|
||
|
|
"name": "stdout",
|
||
|
|
"output_type": "stream",
|
||
|
|
"text": [
|
||
|
|
"Epoch 1/3\n",
|
||
|
|
"1875/1875 [==============================] - 138s 46ms/step - loss: 1.7400 - metric: 0.7341\n",
|
||
|
|
"Epoch 2/3\n",
|
||
|
|
"1875/1875 [==============================] - 88s 47ms/step - loss: 1.5284 - metric: 0.9748\n",
|
||
|
|
"Epoch 3/3\n",
|
||
|
|
"1875/1875 [==============================] - 84s 45ms/step - loss: 1.5035 - metric: 0.9816\n"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"data": {
|
||
|
|
"text/plain": [
|
||
|
|
"<keras.callbacks.History at 0x2457460ad30>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
"execution_count": 109,
|
||
|
|
"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": {},
|
||
|
|
"source": [
|
||
|
|
"<h2 align='center' dir='rtl'>اماده سازی نمونه ها و محاسبه اهمیت هر نمونه نسبت به حضور دیگر نمونه ها</h2>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 9,
|
||
|
|
"id": "2a669978",
|
||
|
|
"metadata": {},
|
||
|
|
"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": 10,
|
||
|
|
"id": "c910c017",
|
||
|
|
"metadata": {},
|
||
|
|
"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([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": 147,
|
||
|
|
"id": "a4761918",
|
||
|
|
"metadata": {},
|
||
|
|
"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": {},
|
||
|
|
"source": [
|
||
|
|
"<h2 align='center' dir='rtl'>مشاهده تاثیر دیگر نمونه ها بر پیش بینی مدل</h2>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"cell_type": "code",
|
||
|
|
"execution_count": 148,
|
||
|
|
"id": "7b17b3c2",
|
||
|
|
"metadata": {},
|
||
|
|
"outputs": [
|
||
|
|
{
|
||
|
|
"data": {
|
||
|
|
"text/plain": [
|
||
|
|
"<matplotlib.legend.Legend at 0x245a13bb040>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
"execution_count": 148,
|
||
|
|
"metadata": {},
|
||
|
|
"output_type": "execute_result"
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"data": {
|
||
|
|
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAx0AAAE/CAYAAAAjYq6HAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAAA7UklEQVR4nO3dd3xV5eHH8e+TDQlhJUDCSoCwSRhhKChEpgPUVlqs26p1gLPun7XWX61V21+lItS6FRVLUUFRHIAoey9ZYYeZMEICZN7n90cAAwRyAzn33Jt83q8XL3Luee65X68B7jfPc84x1loBAAAAgFOC3A4AAAAAoGqjdAAAAABwFKUDAAAAgKMoHQAAAAAcRekAAAAA4ChKBwAAAABHUToAAAAAOIrSgXNijBlhjFljjDlsjNlojLnI7UwAAADwTyFuB0DgMcYMlPRXSb+WtEBSnLuJAAAA4M8MdyRHRRlj5kh6w1r7httZAAAA4P9YXoUKMcYES0qVFGuMSTfGZBhjXjHG1HA7GwAAAPwTpQMV1VBSqKRrJF0kqbOkLpL+x8VMAAAA8GOUDlTU0WO//9Nau8tamyXp75IuczETAAAA/BilAxVirT0gKUMSJwMBAADAK5QOnIu3JI0yxjQwxtSVdL+kz92NBAAAAH/FJXNxLp6VFCNpvaQ8SR9L+rOriQAAAOC3uGQuAAAAAEexvAoAAACAoygdAAAAABxF6QAAAADgKEoHAAAAAEdROgAAAAA46qyXzB0YNJxLWwE4J994/mPczgD3xcTE2ISEBLdjAAB8YPHixVnW2tiy9nGfDgCAYxISErRo0SK3YwAAfMAYs/VM+1heBQAAAMBRlA4AAAAAjqJ0AAAAAHAU53QAAHyqsLBQGRkZysvLczsKXBYREaEmTZooNDTU7SgAHEbpAAD4VEZGhmrVqqWEhAQZw0XOqitrrfbt26eMjAwlJia6HQeAw1heBQDwqby8PNWvX5/CUc0ZY1S/fn1mvIBqgtIBAPA5Cgckvg+A6oTSAQCQJBlj3jTG7DXGrDrDfmOMGW2MSTfGrDDGdPV1Rn+UkJCgrKys8x7j1GuX9sc//lEvvfTSaY/v3LlT11xzjSRp5syZuuKKKyRJkydP1vPPPy9J+vTTT/XTTz+deM4f/vAHffvtt+cTH4HohRekGTNOfmzGjJLHgbOgdAAAjntb0pCz7L9UUtKxX3dIGuuDTNVeUVGR468RHx+viRMnnvb4sGHD9Nhjj0k6vXT86U9/0oABAxzPBj/Tvbv0q1/9XDxmzCjZ7t7d3Vzwe5QOAIAkyVo7S9L+swy5UtK7tsQ8SXWMMXG+SVd5tmzZorZt2+q2225Tx44ddd111+nbb79V7969lZSUpAULFkiS9u/fr6uuukrJycnq1auXVqxYIUnat2+fBg0apC5duuh3v/udrLUnjv3++++rR48e6ty5s373u9+puLj4rFmioqL00EMPqWvXrurfv78yMzMlSf369dMTTzyhvn376uWXX9Z3332nLl26qFOnTrr11luVn59/4hgvvviievTooR49eig9PV2SNGXKFPXs2VNdunTRgAEDtGfPnhPjly9frksuuURJSUn697//feI96dix42n53n77bY0cOVJz5szR5MmT9fDDD6tz587auHGjbr755hNFZfHixerbt6+6deumwYMHa9euXZKk0aNHq3379kpOTtaIESMq9j8K/iktTfr445Ki8Yc/lPz+8ccljwNnQekAAHirsaTtpbYzjj12EmPMHcaYRcaYRcc/RPub9PR03XfffVqxYoXWrl2rDz74QD/++KNeeuklPffcc5Kkp59+Wl26dNGKFSv03HPP6cYbb5QkPfPMM+rTp4+WLl2qYcOGadu2bZKkNWvWaMKECZo9e7aWLVum4OBgjR8//qw5Dh8+rK5du2rJkiXq27evnnnmmRP7Dh48qO+//1733HOPbr75Zk2YMEErV65UUVGRxo79eZIpOjpaCxYs0MiRI3X//fdLkvr06aN58+Zp6dKlGjFihF4otfRlxYoV+uKLLzR37lz96U9/0s6dO8t9vy688EINGzZML774opYtW6aWLVue2FdYWKhRo0Zp4sSJWrx4sW699VY9+eSTkqTnn39eS5cu1YoVKzRu3LhyXwcBIi1Nuusu6dlnS36ncMALXDIXAOCtss76tac9YO1rkl6TpNTU1NP2l/bMlNX6aeehykl3TPv4aD09tMNZxyQmJqpTp06SpA4dOqh///4yxqhTp07asmWLJOnHH3/Uf//7X0nSJZdcon379ik7O1uzZs3SpEmTJEmXX3656tatK0n67rvvtHjxYnU/tszk6NGjatCgwVlzBAUF6de//rUk6frrr9cvfvGLE/uOP75u3TolJiaqdevWkqSbbrpJY8aMOVEwrr322hO/P/DAA5JKLkv861//Wrt27VJBQcFJl6S98sorVaNGDdWoUUNpaWlasGCBOnfufNacZ7Nu3TqtWrVKAwcOlCQVFxcrLq5kAiw5OVnXXXedrrrqKl111VXn/BrwMzNmSGPHSk89VfJ7WhrFA+WidAAAvJUhqWmp7SaSyv8xuR8KDw8/8XVQUNCJ7aCgoBPnUJReNnXc8astlXXVJWutbrrpJv3lL38551yljxsZGXnGHGd6zvGvR40apQcffFDDhg3TzJkz9cc//rHM8WVtV5S1Vh06dNDcuXNP2/fFF19o1qxZmjx5sp599lmtXr1aISF89Ahox8/hOL6kKi2NJVbwCn/yAQDemixppDHmI0k9JWVba3edzwHLm5Fw08UXX6zx48frqaee0syZMxUTE6Po6OgTj//P//yPvvzySx04cECS1L9/f1155ZV64IEH1KBBA+3fv185OTlq3rz5GV/D4/Fo4sSJGjFihD744AP16dPntDFt27bVli1blJ6erlatWum9995T3759T+yfMGGCHnvsMU2YMEEXXHCBJCk7O1uNG5esfHvnnXdOOt5nn32mxx9/XIcPH9bMmTP1/PPPq6CgoNz3o1atWsrJyTnt8TZt2igzM1Nz587VBRdcoMLCQq1fv17t2rXT9u3blZaWpj59+uiDDz5Qbm6u6tSpU+5rwY8tXHhywTh+jsfChZQOnBWlAwAgSTLGfCipn6QYY0yGpKclhUqStXacpKmSLpOULumIpFvcSeobf/zjH3XLLbcoOTlZNWvWPPHh/emnn9a1116rrl27qm/fvmrWrJkkqX379vrf//1fDRo0SB6PR6GhoRozZsxZS0dkZKRWr16tbt26qXbt2powYcJpYyIiIvTWW29p+PDhKioqUvfu3XXnnXee2J+fn6+ePXvK4/Howw8/PJF9+PDhaty4sXr16qXNmzefGN+jRw9dfvnl2rZtm5566inFx8efWFJ2NiNGjNDtt9+u0aNHn3Slq7CwME2cOFH33nuvsrOzVVRUpPvvv1+tW7fW9ddfr+zsbFlr9cADD1A4qoJHHjn9MZZXwQvmbNO2A4OGn31OFwDO4BvPf7jrF5SammoXLVp00mNr1qxRu3btXErkX6KiopSbm+t2DFfx/QBUHcaYxdba1LL2cfUqAAAAAI6idAAA4JLqPssBoPqgdAAAAABwFKUDAAAAgKMoHQAAAAAcRekAAAAA4ChKBwAAPnLw4EG9+uqrZ9wfFRV11udv2bJFHTt2rNBr3nzzzSfdVwMA3EDpAADAS0VFRWfdLk95pQMAqipKBwDAf73wgjRjxsmPzZhR8vh5ePfdd5WcnKyUlBTdcMMNkqStW7eqf//+Sk5OVv/+/bVt2zZJJTMFDz74oNLS0vToo4+etr1x40YNGTJE3bp100UXXaS1a9dKkvbs2aOrr75aKSkpSklJ0Zw5c/TYY49p48aN6ty5sx5++OEz5svNzVX//v3VtWtXderUSZ999tmJfUVFRbrpppuUnJysa665RkeOHJEkLV68WH379lW3bt00ePBg7dq167zeIwCoTCFuBwAA4Iy6d5d+9Svp44+ltLSSwnF8+xytXr1af/7znzV79mzFxMRo//79kqSRI0fqxhtv1E033aQ333xT9957rz7
|
||
|
|
"text/plain": [
|
||
|
|
"<Figure size 1080x360 with 2 Axes>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
"metadata": {
|
||
|
|
"needs_background": "light"
|
||
|
|
},
|
||
|
|
"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": 149,
|
||
|
|
"id": "e1124039",
|
||
|
|
"metadata": {},
|
||
|
|
"outputs": [
|
||
|
|
{
|
||
|
|
"data": {
|
||
|
|
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAXQAAAD4CAYAAAD8Zh1EAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAAAcyUlEQVR4nO3df4zb933f8eeb5P3QHWnJ0p14iiRbcnRk683Nkqlu5qJrtiyrnRY1CmyY3a3BggWGgaTLhgGLN2DbH/1jGLoVXdG0hpFmRbGuRpEaq1doSbFua/8o0llpsySOR54i/5KV451+k3e6HyTf+4P8niiap6Mkkt8ffD0AQ8cvv0e+Qete+tzn+/5+PubuiIhI/KXCLkBERAZDgS4ikhAKdBGRhFCgi4gkhAJdRCQhMmG98dzcnJ84cSKstxcRiaVvfOMbl9x9vtdzoQX6iRMnOHv2bFhvLyISS2b29m7PacpFRCQhFOgiIgmhQBcRSQgFuohIQuwZ6Gb2ZTNbMbPv7PK8mdmvmNk5M/uWmX1k8GWKiMhe+hmh/ybw5B2efwpYbP/3HPDr91+WiIjcrT0D3d3/BLhyh1OeBn7LW74OHDCzI4MqUERE+jOIOfSjwLsdjy+0j72PmT1nZmfN7Ozq6uoA3lpEdnN1bYvf/+Z7YZchIzSIQLcex3ousu7uL7n7aXc/PT/f80YnERmQ//J/3uHzL3+Td6+sh12KjMggAv0CcLzj8THg4gBeV0TuQ2m5etufknyDCPRXgU+1u10+Clx39+8P4HVF5D6UK+1AryjQx8Wea7mY2e8AHwPmzOwC8G+ACQB3fxE4A3wSOAesA58eVrEi0p/tRpPvrdaAW8EuybdnoLv7s3s878BnB1aRiNy3ty6tsd1w0inTlMsY0Z2iIgkUTLM88cFDnF9do95ohlyRjIICXSSBypUaKYNPPnaErUaTty6r02UcKNBFEqi8XOXE3CyPHd3feqx59LGgQBdJoHKlSjGf49ThLClT6+K4UKCLJMzGdoO3Lq+xmM8xPZHm4UOzGqGPCQW6SMKcW6nRdCjmcwAU8lkF+phQoIskTBDexYVs6898jrcur7Ox3QizLBkBBbpIwpQrNSbTKR4+NAtAYSFHo+mcX10LuTIZNgW6SMKUK1UemZ9lIt368S60p1407ZJ8CnSRhCktVyku5HYenzg0y0TatKbLGFCgiyRIdWOb967d3BmVA0xmUjwyl2VJgZ54CnSRBFlaaS3I1Rno0JpH1wg9+RToIglSbt9AVOwK9GI+y7tXbrK2WQ+jLBkRBbpIgpQqVfZNpDn24L7bji+2Az4YwUsyKdBFEmSpUqOQz5JK3b4zZFGdLmNBgS6SIKVK9X3z5wDHD84wPZHamZKRZFKgiyTElbUtVqubPQM9nTIWD+vCaNIp0EUSIphOKSy8P9ABFrWmS+Ip0EUSYmcNlx4j9OB45cYm19e3R1mWjJACXSQhSstVHpjOkH9gqufzwci9vKJRelIp0EUSYqlSo7iQw8x6Ph+M3LXZRXIp0EUSwN137XAJHNk/TW4qo3n0BFOgiyTASnWT6ze37xjoZsZiPqsReoIp0EUSIAjpOwU6QHEhR7lSxd1HUZaMmAJdJAF2Whbz2TueV8jnuLq+zaXa1ijKkhFToIskQGm5ylx2ikPZ3h0uAS0BkGwKdJEEKK/UdvYQvZOgdVHz6MmkQBeJuWbTWdqjwyUwl53i4OykRugJpUAXibn3rt1kfavRV6BDa55da7okkwJdJOb67XAJFPM5lio1dbokkAJdJOZKfXa4BAoLOWqbdS5e3xhmWRKCvgLdzJ40s5KZnTOzF3o8v9/M/puZ/V8ze93MPj34UkWkl3KlytED+8hNT/R1/k6niy6MJs6egW5maeCLwFPAo8CzZvZo12mfBb7r7h8CPgb8BzObHHCtItJDabna9+gcbm1Hp3n05OlnhP44cM7dz7v7FvAy8HTXOQ7krLUqUBa4Amg3WpEhqzeanF9d23UN9F7275tg4YFpjdATqJ9APwq82/H4QvtYp18FfhC4CHwb+Ly7N7tfyMyeM7OzZnZ2dXX1HksWkcBbl9fZajR3XQN9N4WFnJbRTaB+Ar3XWpzdl8d/Avgm8AHgrwC/amYPvO+b3F9y99Pufnp+fv4uSxWRbrdu+b+7QC/msyxVajSa6nRJkn4C/QJwvOPxMVoj8U6fBl7xlnPAm8APDKZEEdlNabmKGZw63P8cOrT+AdisN3nnyvqQKpMw9BPorwGLZnayfaHzGeDVrnPeAT4OYGZ5oAicH2ShIvJ+5UqVE4dmmZ5I39X3FbTZRSLtGejuXgc+B3wNeAP4XXd/3cyeN7Pn26f9AvCEmX0b+CPgC+5+aVhFi0hLa1OLuxudQ2vDaNAiXUmT6eckdz8DnOk69mLH1xeBvz3Y0kTkTja2G7x9eZ2feuzIXX/vzGSGhw7OKNATRneKisTU+dU1Gk2/q5bFToV8ToGeMAp0kZi61w6XQHEhy/nVNbbq7+swlphSoIvEVKlSZSJtnDg0e0/fX8jnqDedNy+tDbgyCYsCXSSmystVHpnLMpm5tx/jgpYASBwFukhMlSrVe54/B3hkfpZ0ylhSoCeGAl0khtY261y4epPiPbQsBqYyaU7OzaoXPUEU6CIxtLRSA+79gmigqE6XRFGgi8RQ+S53KdrNYj7L21fWubnVGERZEjIFukgMlSpVpidSHD84c1+vU8zncIdz7RG/xJsCXSSGypUqi4dzpFO9FkPtX3BRVdMuyaBAF4mh1i5F9zfdAvDwwRkmMykFekIo0EVi5tr6FivVTYoL997hEsikU5yaz6oXPSEU6CIxU64MpsMlUMhntR1dQijQRWKmdJ9ruHQrLOS4eH2DGxvbA3k9CY8CXSRmystVclMZjuyfHsjrBfuRLlXU6RJ3CnSRmAlu+Te7vw6XQDDS14XR+FOgi8SIu1OuDKbDJXD0wD5mJ9NaAiABFOgiMbJa3eTa+vZ9reHSLZUyTmkJgERQoIvEyE6Hy32ssthLMZ9VoCeAAl0kRgbd4RIo5HNcqm1xubY50NeV0VKgi8RIebnKodlJ5rJTA33d4s4SAOp0iTMFukiMlAZ8QTSgTpdkUKCLxESz6SxVqjuj6UE6nJti/74JLQEQcwp0kZh479pN1rYaQxmhmxnFfE7b0cWcAl0kJpZWWmE7iEW5eiksZCktV3H3oby+DJ8CXSQmSsutC5anDg9+hA6tJQBubNSp3FCnS1wp0EViolypcmT/NPv3TQzl9RfbUzmaR48vBbpITAxqU4vd7HS6aAmA2FKgi8RAvdHk3GptKB0ugYOzk8znptS6GGMKdJEYePvKOlv15lBH6NCaR1egx5cCXSQGgnbC4pADvZDPUa7UaDbV6RJHfQW6mT1pZiUzO2dmL+xyzsfM7Jtm9rqZ/fFgyxQZb6XlGmZw6vBwWhYDhXyWm9sNLly9OdT3keHYM9DNLA18EXgKeBR41swe7TrnAPBrwE+7+18C/u7gSxUZX+VKlYcOzrBvMj3U9wlWcVSnSzz1M0J/HDjn7ufdfQt4GXi665yfBV5x93cA3H1lsGWKjLdhreHSbbH9G4Dm0eOpn0A/Crzb8fhC+1inAvCgmf1vM/uGmX2q1wuZ2XNmdtbMzq6urt5bxSJjZrPe4M1La0OfPwfITU9w9MA+BXpM9RPovTYu7L5ikgH+KvCTwE8A/8rMCu/7JveX3P20u5+en5+/62JFxtH51TUaTR/4pha7KS7ktB1dTPUT6BeA4x2PjwEXe5zzVXdfc/dLwJ8AHxpMiSLjrTyiDpfAYj7L+dU1thvNkbyfDE4/gf4asGhmJ81sEngGeLXrnN8HfszMMmY2A/wI8MZgSxUZT+VKlUzKODk3O5L3K+ZzbDWavH15bSTvJ4OT2esEd6+b2eeArwFp4Mvu/rqZPd9+/kV3f8PMvgp8C2gCX3L37wyzcJFxUVqucXJulsnMaG4bubXZRW1oC4HJcOwZ6ADufgY403Xsxa7Hvwj84uBKExFojdAfO7Z/ZO936nCWlLXWjvnkY0dG9r5y/3SnqEiErW/
|
||
|
|
"text/plain": [
|
||
|
|
"<Figure size 432x288 with 1 Axes>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
"metadata": {
|
||
|
|
"needs_background": "light"
|
||
|
|
},
|
||
|
|
"output_type": "display_data"
|
||
|
|
},
|
||
|
|
{
|
||
|
|
"data": {
|
||
|
|
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA0cAAALxCAYAAACXXt1fAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjQuMywgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy/MnkTPAAAACXBIWXMAAAsTAAALEwEAmpwYAAEAAElEQVR4nOz9d5Rd2XnYif723ifdHCpHFDIa3ejIZjOKFEVSokhJli3bcpJkOYxHluVnz2iel98bh3GQJb9xnLHH47EtB0m2cjIpihRzs5tks3M3uhupUFWoXDeHc0/Y+/1xbqELaACNjgAa57dWLeCefPd39z77+/YXhDGGlJSUlJSUlJSUlJSU2x15ox8gJSUlJSUlJSUlJSXlZiBVjlJSUlJSUlJSUlJSUkiVo5SUlJSUlJSUlJSUFCBVjlJSUlJSUlJSUlJSUoBUOUpJSUlJSUlJSUlJSQFS5SglJSUlJSUlJSUlJQVIlaOUlJSUlJSUlJSUlBTgTVKOhBB/TQixLoRoCiH+vRDCvcax9wohvi2E6A3/vXfPvruEEJ8VQmwLIV5RgEkIcYcQ4gvD+5wWQvzgVe7xt4UQRgjx0T3bvlMI8cXhuYtXOOeLQogtIURLCPGUEOIH9uybEkL8thBidXjdhetunJuMm0FWQojjQojHhBD14d/nhRDHr3ANRwjxghBi5SrP96GhPP7+nm3fKYR4RgjREELsCCF+Qwgx8xqa6KbgVpCTEOKnhRDPCiHaQohzQoifvuzaf28oi0gI8Xcu2/dhIYQWQnT2/P3o62utG8tNIqv3CCE+J4SoDcexXxFCTF3hGlfsU9eS1XD/XxnKuDX8TXzg+lvo5uFtlNWCEOLTw36zLoT4P4QQ1p79f34ow44Q4veEENN79gkhxM8Ox68dIcTPCSHEcN+4EOKXRPIuagohHhZCPHSV5/8PIhkfD73O5rqh3AyyEkL8qcvGqN6wTR8Y7i8LIf6jEGJz+Pd3Lrv2ohCiv+f839+zLx0DbyJZ7bnHleYV6RzwzZXVwrAd98rrf91z7mcu2xcIIZ657Pp/VSTvpK4Q4qQQ4shw+9+87Nz+sJ+NvqaGMsa8oT/gu4EN4E6gAnwJ+EdXOdYBzgN/DXCBnxp+dob7jwJ/DviB5NEuOdcCXgL+OqCAjwBd4Mhlxx0EngFWgY/u2f5u4M8AfxFYvMKz3Q1Yw/8/BLSBqeHnCeAngPcCBlh4o+12I/5uFlkBZWABEMP9PwU8fYVn+P8AXwFWrrDPBp4EHgX+/p7tE8D08P8u8HPAb9/otn8nygn4X4D7h9c5OrzvD+/Z/6PAJ4DfAv7OZff+8JXkeqv93USy+gTwR4EikAX+PfB7V3iGK/apV5HVQ8N7PTD8LfyPwBagbnT734yyGu7/NPDzgAdMkryTfmq470PA5vA5HOBfA1/ec+7/ALwIzAIzwPPAXxruOzD8DUwNfwd/EdgG8pfd/wNDORvg0I1u+1tVVlc49seAM4AYfv4PwK8M+9zCcN+f3XP8InvmIZdd68OX98Nb8e+dIqvhMdeaV6RzwDdvDFwYtqN1nc/9JeBv7fn854GngeMk76SDQPUq5/4d4Auvua3ehMb+ReAf7vn8XcD6VY79OHBh98c63LYEfM9lxx26vLGBu4DOZef+PvD3LjvuM8D3cpVBCfgoV1COLjvm3YAPvPuy7dYt3jFuKlntadO/DPQu274fOEkyYbuScvQ3SBSfn2fPIHbZMS7wM8DzN7rt36lyuuyYfwH8yyts/y+8c5Wjm05Ww333A+3Ltl2zT11DVn8c+Oaez7nhODh1o9v/ZpTVcPtJ4Hv3fP7HwL8Z/v//B/yfe/ZND9vz4PDz14G/uGf/nwMevcb3agEP7PlsAU+QGPxuVeXoppDVFY79IvC393zeBh7c8/lvAl/d83mRd75y9I6Q1XDbNecVpHPAN2sMXOA6laPhsTGwf/hZAsvAd13HuYJECf7R19pWb4Zb3Z3AU3s+PwVMCCFGrnLs02b41EOeHm5/NcRVtt118YMQfxQIjDGfvo7rvfJiQvyuEMIHvkGiqT72eq5zE3PTyApACNEgUUL/JfAPLzv+X5IMXv1XXEiIfcCPA//bFW8uxPzw2n3gfyYZ7G4lbiU57R4jgA8Cz13HfXcZF0JsDJfG/6kQIvcazr1ZuKlktYfv4JWyuGqfehU+AyghxENCCEXS954E1l/jdW40b5esAP458MNCiKxI3Ho/AfzecJ/gUnnu/n9Xlld6zived+jm4gCn92z+a8BXjDFPX+ez3ozcLLK6yPC98x3Af7p812X/v7xP/oJIXF1/Xwhxz2X70jHwJpHVq80r3iHcbLI6L4RYEYkL8NXc3n6ERIk9N/w8O/y7SwixPOw7f1cIcSV95oMkq36/dp3PfJE3QznKA809n3f/X7iOY3ePv9Kxl/MCiSvCTwshbCHEx0ncE7IAQog8ycTt/3XdT34ZxphPDZ/le4HPGmP0673WTcpNIatdjDFloAT8JImlEwCRxFJYxpjfuMr1/wXwvxpjOlfaaYxZGl57FPj/Dp/nVuKWkNNl/B2S8eQ/XMd9d+99L4l70EdIXLb+yXWeezNxU8kKQAhxN/C3gJ/es+3V+tS1aJO8XL4GDIC/TbKyYa551s3H2yUrgC+TTCJawAqJoe03h/s+DfwxIcTdQogMiawML8vySs+ZHxogLiKEKAL/Gfi7xpjmcNsciVve37rO57xZuVlktZfLJ2mQTPb+hhCiIJLYrh/n0j75p0gs3/tIVjI+K4QoD/elY+DLx98MsrrmvOIdws0iq23gQZJ+8cDwmr9wlev8CMlK3i6zw38/DpwAvhP4EyQr7Jfzo8Cvvh6ZvmblSFwa9PYZEleP4p5Ddv/fvsLplx+7e/yVjr0EY0wI/CHgkyQWy/8J+GWSRgf4u8B/vqwzvGaMMaEx5jPAdwshvv+NXOtGcxPLau+xXeD/Av6TSIKNcyQrPX/lKt/p+4CCMea/Xcdz1ID/CPyW2BMMfbNxK8rpsuf/SZIB7JPGmMGr3Xd4vXVjzPPGGD3ss/8L8EPXc+6N5GaX1fCl/xngrxpjvjrcds0+dR38eZKJxG6MzJ8GflfsSSJwM3KjZDW0YH4W+HUSF8RREv/+nwUwxvwBiYL5ayQ+/IvD6+7K8krP2dmrjA6Vqt8hcbf7mT3H/jPgf9tVlm4VblZZXcaPkLxP9vJTJCuxp0ji9X6JPX3SGPOwMaZvjOkN5dQgsWanY+Clx99QWb2WecWtxM0qK2NMxxjzmDEmMsZskBhePz40+Oy9zgdI4pV+dc/mXc+HnzPGNIwxi8C/IVnU2HtuhiQO9/LfwfXxWv3wruDT94vAP9jz+SNc24dxhUt9GM9znT6MV7je14H/Yfj/J0m00fXhXwzUgP/3Zee8aszR8LjPA3/tsm3vBH/TGy6rK+yzSH7w95FY0sI9cqwNZblOYoH7ZySWiN39fZJO/FtXufbsUGZXDNa7Gf9uBTnt2fbjw/sfuMY1XxHHcoVjHgJqN7rtb2VZkVjhFhkG7+/Zfs0+9WqyAv4P4J9etu1J4IdudPvfjLIimQgYoLRn2x8Cnr3KvY6QJLyo7JHrX9iz/8fZE3NEEkv52eH3kZddq0EScL0ra0OSPONP3uj2v5VlBbx/KKPCqzz3PwR+6Rr7TwLff5V96Rh4g2TFdc4rSOeAb9UYOHH58cPt/xb4T5dty5J4MHzHnm3/E/Ablx33p0jeh+JK93zVtnoTGvt7hj+m4ySa4Rd49ewXf5VkgP9JLs1+IUgyWxwfNpQHuHvOv3u4LUsSS3Judz8wQqJh7v4tk2iN+eF+OTz3E8N7envue2y4PUOSreRPAwFw/557e7wciHwU8G70D/0WltXHSBQhRWKJ+Bck2QU9ksFnrxz/8HDf5PD4wmX7/xvwTxkqP8Pjjw7lPUZiXX/8Rrf9O01Ow/1/avicd1zl2ezhtX8R+PvD/6vhvg8D88PnmyNxOfkPN7rtb2FZzZAEnv70Fe57zT51HbL6UZJMeQeGz/gxoAccu9HtfxPL6ixJcLdFkvXxN4BfGO7zSGIdxLAPfIl
|
||
|
|
"text/plain": [
|
||
|
|
"<Figure size 1080x1080 with 64 Axes>"
|
||
|
|
]
|
||
|
|
},
|
||
|
|
"metadata": {
|
||
|
|
"needs_background": "light"
|
||
|
|
},
|
||
|
|
"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": {
|
||
|
|
"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
|
||
|
|
}
|