{
  "cells": [
    {
      "cell_type": "markdown",
      "metadata": {
        "_cell_guid": "bb188cd2-27aa-ad75-2079-3cb350d33588"
      },
      "source": [
        "# Classification of  MNIST dataset unsing k-means pre-training"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {
        "_cell_guid": "73c232a1-a823-bfcc-069d-7c9848d40f84"
      },
      "source": [
        "## Importing Modules"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "22b79f82-4296-3dc2-f70c-4b99b4cf808a"
      },
      "outputs": [],
      "source": [
        "import random\n",
        "import numpy as np\n",
        "import matplotlib.pyplot as plt\n",
        "\n",
        "from sklearn.cluster import MiniBatchKMeans\n",
        "from sklearn.linear_model import Perceptron"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {
        "_cell_guid": "335637be-7851-7b4e-a1b6-845ac30825f0"
      },
      "source": [
        "## Importing Data"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "8a22f411-3498-4f71-b158-08d4acd1fd77"
      },
      "outputs": [],
      "source": [
        "#import the dataset\n",
        "print(\"imposting dataset...\")\n",
        "f = open(\"../input/train.csv\")\n",
        "L = f.read().split('\\n')[1:-1]\n",
        "f.close()\n",
        "\n",
        "L = list(map(lambda txt: txt.split(','), L))\n",
        "L = [list(map(int, lst)) for lst in L]\n",
        "\n",
        "#shuffle the dataset\n",
        "print(\"shuffling dataset...\")\n",
        "random.shuffle(L)\n",
        "\n",
        "print(\"splitting dataset...\")\n",
        "\n",
        "#spliting the dataset into a training and testing set\n",
        "training_set = L[:int(.8 * len(L))]\n",
        "testing_set  = L[int(.8 * len(L)):]\n",
        "\n",
        "#separing data and labels...\n",
        "training_set_data   = [L[1:] for L in training_set]\n",
        "training_set_labels = [L[0]  for L in training_set]\n",
        "\n",
        "testing_set_data   = [L[1:] for L in testing_set]\n",
        "testing_set_labels = [L[0]  for L in testing_set]\n",
        "\n",
        "#make the data explotable by reformating them\n",
        "training_set_data = map(lambda lst: np.array(lst).reshape((28, 28)), training_set_data)\n",
        "testing_set_data  = map(lambda lst: np.array(lst).reshape((28, 28)), testing_set_data)\n",
        "\n",
        "testing_set_data  = list(testing_set_data)\n",
        "training_set_data = list(training_set_data)\n",
        "\n",
        "print(\"done.\")"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {
        "_cell_guid": "aab57ceb-f38d-d10a-98ce-40d7e82b9bd4"
      },
      "source": [
        "## Function to extract patches from image"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "6f3f468f-659e-d518-468d-6d3028e20894"
      },
      "outputs": [],
      "source": [
        "def extract_patch(patch_size, image):\n",
        "    L = []\n",
        "    for i in range(len(image) - patch_size):\n",
        "        for j in range(len(image[0]) - patch_size):\n",
        "            L.append(image[i:i+patch_size,j:j+patch_size])\n",
        "    return L"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "ff71bf28-bf0a-f257-d1e7-0b722d720fdd"
      },
      "outputs": [],
      "source": [
        "training_patches_8x8 = sum(map(lambda img: extract_patch(8, img), training_set_data[:1000]), [])\n",
        "training_patches_vector_64 = list(map(lambda img: img.reshape(1, 64)[0], training_patches_8x8))"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "5f6d8787-a89c-7bce-d03b-4e16634db751"
      },
      "outputs": [],
      "source": [
        "#let clusturing the patches \n",
        "mb_kmeans = MiniBatchKMeans(n_clusters=64, max_iter=100) #64 clusters \n",
        "mb_kmeans.fit(training_patches_vector_64)"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {
        "_cell_guid": "f53e7c95-dc69-6fa4-3b91-81a266f9c09f"
      },
      "source": [
        "## Plot the extracted clusters centers"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "8cd0c01d-e0ae-6453-a228-2b40394f9edd"
      },
      "outputs": [],
      "source": [
        "plt.figure(figsize=(4.2, 4))\n",
        "for i, patch in enumerate(mb_kmeans.cluster_centers_):\n",
        "    plt.subplot(8, 8, i + 1)\n",
        "    plt.imshow(patch.reshape((8, 8)), cmap=plt.cm.gray, interpolation='nearest')\n",
        "    plt.xticks(())\n",
        "    plt.yticks(())\n",
        "\n",
        "\n",
        "plt.suptitle('Cluster centers of the patches\\n')\n",
        "plt.subplots_adjust(0.08, 0.02, 0.92, 0.85, 0.08, 0.23)\n",
        "\n",
        "plt.show()"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {
        "_cell_guid": "507efc16-03f6-f041-418d-52113be69e65"
      },
      "source": [
        "## Transorm an image to a set of patches"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "0309a1e6-6333-bbca-bf43-43d8b5197de1"
      },
      "outputs": [],
      "source": [
        "def one_hot(v, max_value):\n",
        "    L = [0] * max_value\n",
        "    L[v] = 1\n",
        "    return L"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "cc42f841-fa07-7dd6-6a4a-cd977423ad95"
      },
      "outputs": [],
      "source": [
        "def transform_img(img):\n",
        "    #make the image a 24x24 pixels image (to be divisible by 8)\n",
        "    img = img[2:-2,2:-2]\n",
        "    \n",
        "    #extract patches from image\n",
        "    patches = list(map(lambda img: img.reshape(1, 64)[0], extract_patch(8, img)))\n",
        "    clusters_assigments = mb_kmeans.predict(patches)\n",
        "    \n",
        "    #return the one hot represenation of the image\n",
        "    return  sum(map(lambda v: one_hot(v, 64), clusters_assigments), [])"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "7dee7809-86f9-8a8d-6ff5-30e38d6a546e"
      },
      "outputs": [],
      "source": [
        "#transform only 4000 images the training set...\n",
        "training_set_data_one_hot = map(transform_img, training_set_data[1001:5001])\n",
        "training_set_data_one_hot = list(training_set_data_one_hot)"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "f5efa984-62df-9cf3-8adb-0bd5b51d84b4"
      },
      "outputs": [],
      "source": [
        "#transform only 1000 images the testing set...\n",
        "testing_set_data_one_hot = map(transform_img, testing_set_data[:1000])\n",
        "testing_set_data_one_hot = list(testing_set_data_one_hot)"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {
        "_cell_guid": "a21afdd1-9dfb-66b5-1b16-a01c63bbdb47"
      },
      "source": [
        "## Training a perceptron classifier on the original dataset"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "75a4a259-ef48-338b-75f3-50fb5fc15051"
      },
      "outputs": [],
      "source": [
        "#train with only 4000 images\n",
        "training_set_data_flattern = map(lambda img: img.reshape(1, 28 ** 2)[0], training_set_data[1001:5001])\n",
        "training_set_data_flattern = list(training_set_data_flattern)"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "9c3f068f-1fd4-c0d3-4b56-dc1ea5ebfbf6"
      },
      "outputs": [],
      "source": [
        "testing_set_data_flattern = map(lambda img: img.reshape(1, 28 ** 2)[0], testing_set_data[:1000])\n",
        "testing_set_data_flattern = list(testing_set_data_flattern)"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "bf695fb6-a222-5b4b-028d-746105968add"
      },
      "outputs": [],
      "source": [
        "original_perceptron = Perceptron()\n",
        "original_perceptron.fit(training_set_data_flattern, training_set_labels[1001:5001])"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "db126275-b67f-4773-9d78-e6ecd893b3ae"
      },
      "outputs": [],
      "source": [
        "print(\n",
        "    \"original perceptron accuracy : \", \n",
        "    original_perceptron.score(testing_set_data_flattern, testing_set_labels[:1000])\n",
        ")"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {
        "_cell_guid": "346d9ec1-4b7d-f5fc-5001-bfc7ed684d4f"
      },
      "source": [
        "## Training a perceptron classifier on the new training dataset"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "e0e24d41-6055-d27f-e799-7c984372feda"
      },
      "outputs": [],
      "source": [
        "perceptron = Perceptron()\n",
        "perceptron.fit(training_set_data_one_hot, training_set_labels[1001:5001])"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "cbcd2ac1-751c-abe6-7dc6-e371e717defa"
      },
      "outputs": [],
      "source": [
        "print(\n",
        "    \"new perceptron accuracy : \", \n",
        "    perceptron.score(testing_set_data_one_hot, testing_set_labels[:1000])\n",
        ")"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {
        "_cell_guid": "83f26be3-048b-9ba0-8b37-8e6ff30e4ef6"
      },
      "source": [
        "## Trying to make something near to pooling\n"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "2b6b8395-a0d9-2595-f433-1e08db1146d9"
      },
      "outputs": [],
      "source": [
        "training_patches_6x6 = sum(map(lambda img: extract_patch(6, img), training_set_data[:1000]), [])\n",
        "training_patches_vector_36 = list(map(lambda img: img.reshape(1, 36)[0], training_patches_6x6))"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "714b9467-c83e-faeb-1acb-895c23703eba"
      },
      "outputs": [],
      "source": [
        "#let clusturing the patches \n",
        "mb_kmeans_6x6_patches = MiniBatchKMeans(n_clusters=36, max_iter=100) #36 clusters \n",
        "mb_kmeans_6x6_patches.fit(training_patches_vector_36)"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "8c73468a-b746-5f27-ff7d-af583518a6ed"
      },
      "outputs": [],
      "source": [
        "plt.figure(figsize=(4.2, 4))\n",
        "for i, patch in enumerate(mb_kmeans_6x6_patches.cluster_centers_):\n",
        "    plt.subplot(6, 6, i + 1)\n",
        "    plt.imshow(patch.reshape((6, 6)), cmap=plt.cm.gray, interpolation='nearest')\n",
        "    plt.xticks(())\n",
        "    plt.yticks(())\n",
        "\n",
        "\n",
        "plt.suptitle('Cluster centers of the 6x6 patches\\n')\n",
        "plt.subplots_adjust(0.08, 0.02, 0.92, 0.85, 0.08, 0.23)\n",
        "\n",
        "plt.show()"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "a73bc33e-6ecf-4eb8-1685-d5cf006fa424"
      },
      "outputs": [],
      "source": [
        "def get_4x4_patches_clusters_of_image(img):\n",
        "    #make the image a 24x24 pixels image (to be divisible by 8)\n",
        "    img = img[2:-2,2:-2]\n",
        "\n",
        "    #extract 4x4 patches of size 8x8\n",
        "    M = [[None] * 4 for _ in range(4)]\n",
        "    for i in range(4):\n",
        "        for j in range(4):\n",
        "            M[i][j] = mb_kmeans_6x6_patches.predict([\n",
        "                img[i*6:(i+1)*6,j*6:(j+1)*6].reshape((1, 36))[0]\n",
        "            ])[0]\n",
        "    return M\n",
        "        "
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {
        "_cell_guid": "7fb9cef4-2bc1-c71f-f641-faa82d91edb5"
      },
      "source": [
        "### Traing perceptron on those mini patches without pooling"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "218f9508-dcf5-eb7a-66c4-6d330531bfe6"
      },
      "outputs": [],
      "source": [
        "#get the patches cluster numbers\n",
        "training_4x4_patches_6x6_set_data = list(map(\n",
        "    lambda img: sum(get_4x4_patches_clusters_of_image(img), []),\n",
        "    training_set_data[1001:5001]\n",
        "))\n",
        "testing_4x4_patches_6x6_set_data = list(map(\n",
        "    lambda img: sum(get_4x4_patches_clusters_of_image(img), []),\n",
        "    testing_set_data[:1000]\n",
        "))"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "254322a2-3820-d13f-b6eb-bd042453a9f2"
      },
      "outputs": [],
      "source": [
        "#apply one hot on the previous data\n",
        "training_4x4_6x6_oh_set_data = map(\n",
        "    lambda lst: sum(map(lambda n: one_hot(n, 36), lst), []),\n",
        "    training_4x4_patches_6x6_set_data\n",
        ")\n",
        "testing_4x4_6x6_oh_set_data = map(\n",
        "    lambda lst: sum(map(lambda n: one_hot(n, 36), lst), []),\n",
        "    testing_4x4_patches_6x6_set_data\n",
        ")"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "0c46e923-0d95-a3db-def4-49a5dca924d4"
      },
      "outputs": [],
      "source": [
        "training_4x4_6x6_set_data = list(training_4x4_6x6_oh_set_data)\n",
        "testing_4x4_6x6_set_data = list(testing_4x4_6x6_oh_set_data)"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "a97b7b8e-b562-27cd-2964-6e0424b3bd90"
      },
      "outputs": [],
      "source": [
        "perceptron_4x4_6x6_no_pooling = Perceptron()\n",
        "perceptron_4x4_6x6_no_pooling.fit(training_4x4_6x6_set_data, training_set_labels[1001:5001])"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "d01fdd9b-37aa-1b0f-d6a1-b14e20f982ad"
      },
      "outputs": [],
      "source": [
        "print(\n",
        "    \"perceptron 4x4 6x6 no pooling accuracy : \", \n",
        "    perceptron_4x4_6x6_no_pooling.score(testing_4x4_6x6_set_data, testing_set_labels[:1000])\n",
        ")"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {
        "_cell_guid": "2d821ff2-2048-24b1-f8c4-51a87eb20244"
      },
      "source": [
        "### Now try to train a perceptron with some pooling applied on the data before"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "7253f3ff-5b30-29d4-108e-c49c9f24a5a8"
      },
      "outputs": [],
      "source": [
        "print(training_4x4_patches_6x6_set_data[0])"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "0b520440-5fcd-4ff3-0a18-9aec30b17ee9"
      },
      "outputs": [],
      "source": [
        "def pool_4x4_patches_into_2x2(L):\n",
        "    #make 4 pooling groups\n",
        "    G1 = [ L[0],  L[1],  L[4],  L[5]] #corner left  top\n",
        "    G2 = [ L[2],  L[3],  L[6],  L[7]] #corner right top\n",
        "    G3 = [ L[8],  L[9], L[12], L[13]] #corner left  bottom\n",
        "    G4 = [L[10], L[11], L[14], L[15]] #corner right bottom\n",
        "    #one hot them and regroup by pool\n",
        "    out = [1 if i in G1 else 0 for i in range(36)] + [1 if i in G2 else 0 for i in range(36)] + [1 if i in G3 else 0 for i in range(36)] + [1 if i in G4 else 0 for i in range(36)]\n",
        "    return out\n",
        "\n",
        "print(pool_4x4_patches_into_2x2(training_4x4_patches_6x6_set_data[1]))"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "cbcd1e30-05ef-c9da-76cf-545056089bd5"
      },
      "outputs": [],
      "source": [
        "#make the pooled ttrainingg and testing sets\n",
        "training_pooled_set_data = list(map(pool_4x4_patches_into_2x2, training_4x4_patches_6x6_set_data))\n",
        "testing_pooled_set_data = list(map(pool_4x4_patches_into_2x2,  testing_4x4_patches_6x6_set_data))"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "2b68c0ca-4a87-8189-da20-c33e24dd175d"
      },
      "outputs": [],
      "source": [
        "#train the pooled percpetron\n",
        "perceptron_pooled = Perceptron()\n",
        "perceptron_pooled.fit(training_pooled_set_data, training_set_labels[1001:5001])"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "cf6bc363-eb1d-7e62-96e9-21162ae44a5d"
      },
      "outputs": [],
      "source": [
        "print(\n",
        "    \"pooled perceptron accuracy : \", \n",
        "    perceptron_pooled.score(testing_pooled_set_data, testing_set_labels[:1000])\n",
        ") #hummm mutch lower.... hummm... I missed something... hummm... TODO: make it better... :p"
      ]
    },
    {
      "cell_type": "markdown",
      "metadata": {
        "_cell_guid": "9c213752-8f90-34bf-2a0a-bf26adb0b63a"
      },
      "source": [
        "## Try a BoW approach using the k-meaan extracted features"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "d3888a50-9b36-8eef-71cf-db7832e01c93"
      },
      "outputs": [],
      "source": [
        "def words_extractor_8x8(img):\n",
        "    #make the image a 24x24 pixels image (to be divisible by 8)\n",
        "    img = img[2:-2,2:-2]\n",
        "\n",
        "    #extract patches from image\n",
        "    patches = list(map(lambda img: img.reshape(1, 64)[0], extract_patch(8, img)))\n",
        "    clusters_assigments = mb_kmeans.predict(patches)\n",
        "\n",
        "    #make the BoW\n",
        "    L = [0] * 64\n",
        "    for c in clusters_assigments:\n",
        "        L[c] += 1\n",
        "\n",
        "    return L\n",
        "\n",
        "print(words_extractor_8x8(training_set_data[0]))"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "cf12bf73-0a59-2971-5f9a-04ae6f1f777e"
      },
      "outputs": [],
      "source": [
        "#extract BoW for makint the trainnig and testing set\n",
        "training_bow_8x8_set_data = list(map(words_extractor_8x8, training_set_data[1001:5001]))\n",
        "testing_bow_8x8_set_data = list(map(words_extractor_8x8, testing_set_data[:1000]))"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "8ae26b82-86f7-9358-dece-eb3c4e066ca9"
      },
      "outputs": [],
      "source": [
        "#train a perceptron\n",
        "perceprton_bow_8x8 = Perceptron()\n",
        "perceprton_bow_8x8.fit(training_bow_8x8_set_data, training_set_labels[1001:5001])"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "25358e59-4588-c1b2-cdc6-55dd0d6739fb"
      },
      "outputs": [],
      "source": [
        "#test the perceptron\n",
        "print(\n",
        "    \"BoW perceptron accuracy : \", \n",
        "    perceprton_bow_8x8.score(testing_bow_8x8_set_data, testing_set_labels[:1000])\n",
        ")"
      ]
    },
    {
      "cell_type": "code",
      "execution_count": null,
      "metadata": {
        "_cell_guid": "ba543b87-fa44-0c9e-925e-3362b1b0c31d"
      },
      "outputs": [],
      "source": [
        ""
      ]
    }
  ],
  "metadata": {
    "_change_revision": 0,
    "_is_fork": false,
    "kernelspec": {
      "display_name": "Python 3",
      "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.6.0"
    }
  },
  "nbformat": 4,
  "nbformat_minor": 0
}