{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "intro",
   "metadata": {},
   "source": [
    "# 16　前の情報を次へ渡すRNN\n",
    "\n",
    "FUJIMOTO LAB 深層学習コース。Web教材の図と説明を読んでから実行してください。Pythonの計算はColabの実行環境で行います。\n",
    "\n",
    "**この回の目標**：窓の作り方とRNNの入力形状を説明できる\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",
    "rnn = nn.RNN(input_size=1,hidden_size=4,batch_first=True)\n",
    "x = torch.tensor([[[18.],[19.],[20.]]])\n",
    "out,state = rnn(x)\n",
    "print(out.shape,state.shape)"
   ],
   "execution_count": null,
   "outputs": []
  },
  {
   "cell_type": "markdown",
   "id": "explain-one",
   "metadata": {},
   "source": [
    "**確かめ方**：1件×3時刻の入力に対し、各時刻の状態と最後の状態が得られます。"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "part-two",
   "metadata": {},
   "source": [
    "## 2　窓をずらす\n",
    "\n",
    "3日分を入力し、4日目を正解とします。窓を1日ずらすと次の学習例ができます。"
   ]
  },
  {
   "cell_type": "code",
   "id": "demo-two",
   "metadata": {},
   "source": [
    "values = torch.tensor([18.,19.,20.,21.,22.])\n",
    "windows = torch.stack([values[i:i+3] for i in range(2)])\n",
    "answers = values[3:]\n",
    "print(windows,answers)"
   ],
   "execution_count": null,
   "outputs": []
  },
  {
   "cell_type": "markdown",
   "id": "explain-two",
   "metadata": {},
   "source": [
    "**結果を読む**：未来の値を入力へ先取りせず、次の値だけを答えにします。"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "experiment",
   "metadata": {},
   "source": [
    "## 3　値を変えて比べる\n",
    "\n",
    "次のセルでは、表示される値や条件を変えて、何が結果を決めるかを確かめます。長くかかる実験は、少数の合成データで行います。"
   ]
  },
  {
   "cell_type": "code",
   "id": "lab",
   "metadata": {},
   "source": [
    "# 正弦波から3点を入力し、次の1点を予測\n",
    "torch.manual_seed(3)\n",
    "series = torch.sin(torch.linspace(0,12,80))\n",
    "samples = torch.stack([series[i:i+3] for i in range(77)]).unsqueeze(-1)\n",
    "targets = series[3:].unsqueeze(-1)\n",
    "body = nn.RNN(1,8,batch_first=True)\n",
    "head = nn.Linear(8,1)\n",
    "opt = torch.optim.Adam(list(body.parameters())+list(head.parameters()),lr=0.02)\n",
    "for epoch in range(80):\n",
    "    opt.zero_grad()\n",
    "    states,_ = body(samples[:60])\n",
    "    prediction = head(states[:,-1])\n",
    "    loss = nn.MSELoss()(prediction,targets[:60])\n",
    "    loss.backward(); opt.step()\n",
    "with torch.no_grad():\n",
    "    states,_ = body(samples[60:])\n",
    "    test_loss = nn.MSELoss()(head(states[:,-1]),targets[60:]).item()\n",
    "print('後半を残したテストのMSE:',round(test_loss,4))\n"
   ],
   "execution_count": null,
   "outputs": []
  },
  {
   "cell_type": "markdown",
   "id": "extra-guide-1",
   "metadata": {},
   "source": [
    "## 追加実習 1　窓と答えの対応を確認\n",
    "\n",
    "順番に並んだ値を、入力3時刻と次の1時刻に切り出します。最後の答えを入力に混ぜないことが大切です。\n",
    "\n",
    "**実行前に予想**：何が変わり、何が変わらないでしょうか。"
   ]
  },
  {
   "cell_type": "code",
   "id": "extra-code-1",
   "metadata": {},
   "source": [
    "for index in [0, 1, 2, 60]:\n",
    "    print('入力', torch.round(samples[index,:,0]*100)/100,\n",
    "          '次の値', round(targets[index,0].item(), 3))\n",
    "print('入力の形:', tuple(samples.shape), '答えの形:', tuple(targets.shape))\n"
   ],
   "execution_count": null,
   "outputs": []
  },
  {
   "cell_type": "markdown",
   "id": "extra-reflect-1",
   "metadata": {},
   "source": [
    "**確認**：予想と違った点を一つ書き、値を一つ変えて再実行してください。"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "reflection",
   "metadata": {},
   "source": [
    "## 自分の言葉で答えよう\n",
    "\n",
    "[18,19,20]の次を予測する窓で、答えにする値は？\n",
    "\n",
    "- まず予想を書く\n",
    "- コードのどの行が答えを決めるか指す\n",
    "- 条件や値を1つ変えて、予想と実行結果を比べる\n",
    "\n",
    "**ヒント**：入力の直後の値を答えにします。"
   ]
  }
 ],
 "metadata": {
  "colab": {
   "name": "16-rnn.ipynb",
   "provenance": []
  },
  "kernelspec": {
   "display_name": "Python 3",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
