{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "intro",
   "metadata": {},
   "source": [
    "# 11　CNNで画像を分類する\n",
    "\n",
    "FUJIMOTO LAB 深層学習コース。Web教材の図と説明を読んでから実行してください。Pythonの計算はColabの実行環境で行います。\n",
    "\n",
    "**この回の目標**：畳み込み・プーリング・分類の流れを言える\n",
    "\n",
    "このノートブックは小さな合成データを使用し、元の教材の固定Driveパスや動画を必要としません。"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "predict",
   "metadata": {},
   "source": [
    "## 1　まず予想する\n",
    "\n",
    "コードを実行する前に、表示される値や形を予想してください。"
   ]
  },
  {
   "cell_type": "code",
   "id": "demo-one",
   "metadata": {},
   "source": [
    "import torch\n",
    "from torch import nn\n",
    "net = nn.Sequential(nn.Conv2d(1,4,3,padding=1),nn.ReLU(),nn.MaxPool2d(2),nn.Flatten(),nn.Linear(4*4*4,2))\n",
    "print(net(torch.zeros(2,1,8,8)).shape)"
   ],
   "execution_count": null,
   "outputs": []
  },
  {
   "cell_type": "markdown",
   "id": "explain-one",
   "metadata": {},
   "source": [
    "**確かめ方**：8×8画像2枚が、2種類の点数へ変わります。ここでは学習前です。"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "part-two",
   "metadata": {},
   "source": [
    "## 2　初めて見る画像で試す\n",
    "\n",
    "学習に使った絵だけで判断すると、覚えただけでも高得点に見えます。学習用とテスト用を別々に作ります。"
   ]
  },
  {
   "cell_type": "code",
   "id": "demo-two",
   "metadata": {},
   "source": [
    "prediction = net(torch.zeros(2,1,8,8))\n",
    "print(prediction.argmax(dim=1))"
   ],
   "execution_count": null,
   "outputs": []
  },
  {
   "cell_type": "markdown",
   "id": "explain-two",
   "metadata": {},
   "source": [
    "**結果を読む**：argmaxは大きい点数の位置を返します。正しさは別の画像と正解で確認します。"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "experiment",
   "metadata": {},
   "source": [
    "## 3　値を変えて比べる\n",
    "\n",
    "次のセルでは、表示される値や条件を変えて、何が結果を決めるかを確かめます。長くかかる実験は、少数の合成データで行います。"
   ]
  },
  {
   "cell_type": "code",
   "id": "lab",
   "metadata": {},
   "source": [
    "# 合成した縦線・横線を小さなCNNに学習させる\n",
    "torch.manual_seed(8)\n",
    "images = torch.zeros(160,1,8,8)\n",
    "labels = torch.arange(160)%2\n",
    "for i in range(160):\n",
    "    if labels[i] == 0:\n",
    "        images[i,0,:,3] = 1  # 縦線\n",
    "    else:\n",
    "        images[i,0,3,:] = 1  # 横線\n",
    "images += torch.randn_like(images)*0.08\n",
    "train_x,test_x = images[:120],images[120:]\n",
    "train_y,test_y = labels[:120],labels[120:]\n",
    "model = nn.Sequential(nn.Conv2d(1,4,3,padding=1),nn.ReLU(),nn.MaxPool2d(2),nn.Flatten(),nn.Linear(64,2))\n",
    "optimizer = torch.optim.Adam(model.parameters(),lr=0.02)\n",
    "for epoch in range(20):\n",
    "    model.train(); optimizer.zero_grad()\n",
    "    loss = nn.CrossEntropyLoss()(model(train_x),train_y)\n",
    "    loss.backward(); optimizer.step()\n",
    "model.eval()\n",
    "with torch.no_grad():\n",
    "    accuracy = (model(test_x).argmax(1)==test_y).float().mean().item()\n",
    "print('別に残した画像の正解率:',round(accuracy,3))\n"
   ],
   "execution_count": null,
   "outputs": []
  },
  {
   "cell_type": "markdown",
   "id": "extra-guide-1",
   "metadata": {},
   "source": [
    "## 追加実習 1　テスト用の間違いも見る\n",
    "\n",
    "正解率だけでなく、縦線と横線をそれぞれ何件ずつ取り違えたか数えます。\n",
    "\n",
    "**実行前に予想**：何が変わり、何が変わらないでしょうか。"
   ]
  },
  {
   "cell_type": "code",
   "id": "extra-code-1",
   "metadata": {},
   "source": [
    "with torch.no_grad():\n",
    "    guessed = model(test_x).argmax(dim=1)\n",
    "for correct_name, correct_id in [('縦線', 0), ('横線', 1)]:\n",
    "    mask = test_y == correct_id\n",
    "    print(correct_name, '件数', mask.sum().item(), '正解', (guessed[mask] == test_y[mask]).sum().item())\n"
   ],
   "execution_count": null,
   "outputs": []
  },
  {
   "cell_type": "markdown",
   "id": "extra-reflect-1",
   "metadata": {},
   "source": [
    "**確認**：予想と違った点を一つ書き、値を一つ変えて再実行してください。"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "reflection",
   "metadata": {},
   "source": [
    "## 自分の言葉で答えよう\n",
    "\n",
    "CNNの畳み込み層が主に調べるのは？\n",
    "\n",
    "- まず予想を書く\n",
    "- コードのどの行が答えを決めるか指す\n",
    "- 条件や値を1つ変えて、予想と実行結果を比べる\n",
    "\n",
    "**ヒント**：小さな窓で画素の組み合わせを調べます。"
   ]
  }
 ],
 "metadata": {
  "colab": {
   "name": "11-cnn.ipynb",
   "provenance": []
  },
  "kernelspec": {
   "display_name": "Python 3",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
