{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "id": "e1c76d2a",
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "np.set_printoptions(precision=3, suppress=True)\n",
    "\n",
    "from matplotlib import pyplot as plt"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "71e2a73f",
   "metadata": {},
   "source": [
    "# Hodology, Pt. I: Toy Example\n",
    "\n",
    "## Setting up the environment"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "id": "253b80e2",
   "metadata": {},
   "outputs": [],
   "source": [
    "n_actions = 3\n",
    "n_states = 4\n",
    "actions = np.arange(n_actions)\n",
    "states = np.arange(n_states)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "id": "e388660c",
   "metadata": {},
   "outputs": [],
   "source": [
    "P0 = np.array([[0,1,0,0],[1,0,0,0],[0,0,0,1],[0,0,1,0]])\n",
    "P1 = np.array([[0,0,1,0],[0,0,0,1],[1,0,0,0],[0,1,0,0]])\n",
    "\n",
    "p = 1/np.sqrt(2)\n",
    "P2 = np.array([[1-p,0,0,p],[0,1-p,p,0],[0,p,1-p,0],[p,0,0,1-p]])\n",
    "\n",
    "# P(s_|s, a)\n",
    "P = np.transpose(np.array([P0, P1, P2]), (1,2,0)) "
   ]
  },
  {
   "cell_type": "markdown",
   "id": "2634eed5",
   "metadata": {},
   "source": [
    "## Tabular Q-learning"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "id": "544e8745",
   "metadata": {},
   "outputs": [],
   "source": [
    "n_episodes = 10_000\n",
    "epsilon = 0.1\n",
    "gamma = 1\n",
    "\n",
    "# Q(s, a, g)\n",
    "Q = np.zeros((n_states, n_actions, n_states))\n",
    "N = np.zeros_like(Q, dtype=int)\n",
    "for g in states:\n",
    "    for _ in range(n_episodes):\n",
    "        s = np.random.randint(n_states)\n",
    "        while  s != g:\n",
    "            if np.random.random() < epsilon:\n",
    "                a = np.random.randint(n_actions)\n",
    "            else:\n",
    "                q = Q[s, :, g]\n",
    "                a = np.random.choice(np.flatnonzero(np.isclose(q, q.max())))\n",
    "\n",
    "            new_s = np.random.choice(n_states, p=P[:,s,a])\n",
    "            if new_s == g:\n",
    "                y = -1\n",
    "            else:\n",
    "                y = -1 + gamma*np.max(Q[new_s, :, g]) \n",
    "\n",
    "            N[s,a,g] += 1\n",
    "            alpha = 1/N[s,a,g]\n",
    "            Q[s, a, g] = (1-alpha)*Q[s,a,g] + alpha*y \n",
    "            s = new_s"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "98a11613",
   "metadata": {},
   "source": [
    "### From cost matrix to geometry"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 5,
   "id": "0e8ca925",
   "metadata": {},
   "outputs": [
    {
     "data": {
      "text/plain": [
       "array([[-0.   ,  1.   ,  1.   ,  1.411],\n",
       "       [ 1.   , -0.   ,  1.399,  1.   ],\n",
       "       [ 1.   ,  1.389, -0.   ,  1.   ],\n",
       "       [ 1.416,  1.   ,  1.   , -0.   ]])"
      ]
     },
     "execution_count": 5,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "V = np.max(Q, axis=1); V\n",
    "C = -V; C"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 6,
   "id": "56d7ef8c",
   "metadata": {},
   "outputs": [
    {
     "data": {
      "text/plain": [
       "array([[0.   , 0.99 , 0.994, 1.415],\n",
       "       [0.99 , 0.   , 1.392, 0.994],\n",
       "       [0.994, 1.392, 0.   , 0.99 ],\n",
       "       [1.415, 0.994, 0.99 , 0.   ]])"
      ]
     },
     "execution_count": 6,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "Pi0 = np.eye(n_states) - np.ones((n_states, n_states))/n_states\n",
    "G = (-1/2)*Pi0 @ (C*C) @ Pi0\n",
    "\n",
    "L, V = np.linalg.eigh(G)\n",
    "positive = L > 1e-1\n",
    "x = V[:, positive]*np.sqrt(L[positive])\n",
    "\n",
    "D = np.array([[np.linalg.norm(a-b) for b in x] for a in x]);\n",
    "D"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 7,
   "id": "263cb0e3",
   "metadata": {},
   "outputs": [
    {
     "data": {
      "text/plain": [
       "<matplotlib.collections.PathCollection at 0x75492ca47b60>"
      ]
     },
     "execution_count": 7,
     "metadata": {},
     "output_type": "execute_result"
    },
    {
     "data": {
      "image/png": "iVBORw0KGgoAAAANSUhEUgAAAl4AAAJGCAYAAACKpItwAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjExLjEsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvctoD+AAAAAlwSFlzAAAPYQAAD2EBqD+naQAAL4lJREFUeJzt3X+UVXW9+P/XAM6gMDPA0AgoKeJcBOEaIqigBgZGKGbpFa8Xb2alaatLWOqHspCbhtdaLV0pZP66a+VSM1C6JpJakIQohkog/qBCwBga+eEMIAwM7O8ffjm3c2FgMM57Bno81po/9vvs9znvs5upZ3ufsynKsiwLAAAKrlVzLwAA4B+F8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCJtmnsBhbBz585YvXp1lJaWRlFRUXMvBwA4hGVZFhs3boxu3bpFq1Z7P6d1SIbX6tWro3v37s29DADgH8iqVavi6KOP3us+h2R4lZaWRsQHB6CsrKyZVwMAHMrq6uqie/fuuf7Ym0MyvHZdXiwrKxNeAEASTfl4kw/XAwAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAk0qa5FwDQmB07s1iwfH3UbNwalaVtY1CPTtG6VVFzLwvgQxNeQIs0a0l1THpiaVTXbs2NdS1vGxNH94mRfbs248oAPjyXGoEWZ9aS6rj6wZfzoisiYk3t1rj6wZdj1pLqZloZwN9HeAEtyo6dWUx6Ymlke3hs19ikJ5bGjp172gOgZRNeQIuyYPn63c50/a0sIqprt8aC5evTLQrgABFeQItSs7Hx6Pow+wG0JMILaFEqS9se0P0AWhLhBbQog3p0iq7lbaOxm0YUxQffbhzUo1PKZQEcEMILaFFatyqKiaP7RETsFl+7tieO7uN+XsBBSXgBLc7Ivl1j6tiTo0t5/uXELuVtY+rYk93HCzhouYEq0CKN7Ns1RvTp4s71wCFFeAEtVutWRXF6z4rmXgbAAeNSIwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgEQKHl6/+93v4vzzz4+TTjopLrrooli0aNE+57z//vvxve99L84444w488wzY+rUqZFlWaGXCgBQUAUNr4ULF8bw4cOjb9++8eMf/zgqKirirLPOij//+c+Nzqmvr4+zzz47pk+fHjfeeGPcfvvt8cYbb8S0adMKuVQAgIIrygp4Kumzn/1s1NbWxq9//euIiMiyLHr37h2f+MQn4q677trjnFtvvTVuvfXWWLZsWXzkIx/JjTc0NESbNm2a9Lp1dXVRXl4etbW1UVZW9ve/EQCARuxPdxT0jNfs2bNj5MiRue2ioqIYNWpUzJ49u9E5Dz/8cFxwwQV50RURTY4uAICWqmDhtWnTpnjvvfeia9eueeNdunSJd955p9F5b775ZvTq1Suuvfba6N+/f3zyk5+Me+65J3bu3NnonPr6+qirq8v7AQBoaQoWXjt27IiIiOLi4rzxkpKSaGho2OOcLMti27ZtcfPNN0eHDh3i3nvvjbFjx8b1118fEydObPS1Jk+eHOXl5bmf7t27H7g3AgBwgBQsvEpLS6O4uDjWrl2bN7527dqoqKjY45yioqLo3LlznHbaafGd73wnBgwYEJdddlmMHz8+fvKTnzT6WhMmTIja2trcz6pVqw7oewEAOBAKFl6tWrWKk08+OV544YW88eeffz5OOeWURucNHDgwOnXqlDdWUVERmzdvbvSWEiUlJVFWVpb3AwDQ0hT0w/VXXXVVTJ8+PRYsWBAREU899VT89re/jS9/+cu5fe68884YNGhQbvsrX/lKzJo1K1599dWI+OAM2T333BOf+tSnoqioqJDLBQAoqIJ+VfDyyy+PZcuWxcc//vEoKyuLzZs3x2233Raf/OQnc/usXbs23nrrrdz2qFGj4uabb46hQ4dGWVlZvPvuu/GpT30qfvzjHxdyqQAABVfQ+3jtsnnz5qipqYmuXbtG27Zt8x5bu3ZtbNiwIaqqqvLG6+vro7q6Orp06bLbnH1xHy8AIJX96Y4k4ZWa8AIAUmkxN1AFAOB/CS8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJJAmvTZs2xbJly2LLli37NW/79u2xZMmSWLVqVYFWBgCQTsHD64YbbojOnTvH0KFDo6KiIr7//e/v19x+/frF+PHjC7hCAIA0Chpe999/f9x1110xb968+Mtf/hIzZsyICRMmxFNPPbXPuTNnzoxf/epXcfbZZxdyiQAAyRQ0vO6+++646KKLYsCAARERcc4558TQoUPj7rvv3uu81atXx5VXXhkPPvhgHH744YVcIgBAMgULr507d8arr74ap556at744MGDY+HChXudN3bs2Bg3blz079+/UMsDAEiuTaGeeOPGjbFt27aoqKjIG6+oqIi1a9c2Ou+WW26JiIivf/3rTX6t+vr6qK+vz23X1dXt52oBAAqvYGe82rT5oOm2bduWN15fXx+HHXbYHucsWrQoJk+eHNdff30sXbo0lixZEhs3boy6urpYsmRJbN++fY/zJk+eHOXl5bmf7t27H9g3AwBwABTsjFe7du2iY8eOUV1dnTdeXV3daBitW7cujjvuuPjGN76RG1u5cmUUFRXFJZdcEk8//XR069Ztt3kTJkyIa6+9NrddV1cnvgCAFqcoy7KsUE9+4YUXxvr162P27NkREZFlWZxwwgkxYsSIuPPOOyMioqamJtatWxe9e/fe43Ocd9550bZt25g2bVqTX7euri7Ky8ujtrY2ysrK/v43AgDQiP3pjoJ+q/HGG2+M+fPnx/XXXx9z586NL33pS/HXv/4174zWlClT4vTTTy/kMgAAWoSChlf//v1j9uzZsWzZsvja174WmzZtirlz58axxx6b26eysjL69OnT6HMcc8wx8dGPfrSQywQASKKglxqbi0uNAEAqLeZSIwAA/0t4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAk0qbQL7Bly5Z4/PHHY8WKFVFVVRUXXHBBtGmz95ddunRpPPfcc7F9+/YYOHBgnHbaaYVeJgBAwRX0jNeGDRti0KBBceutt8aaNWtiwoQJMXTo0Ni6dWujcy699NIYM2ZMvPrqq/HGG2/EyJEj4/LLLy/kMgEAkijKsiwr1JNfd911MX369PjDH/4Q7du3j5qamujVq1d85zvfifHjx+9xzpw5c2Lo0KG57ZdffjkGDBgQzzzzTAwfPrxJr1tXVxfl5eVRW1sbZWVlB+KtAADs0f50R0HPeE2fPj0uvvjiaN++fUREVFZWxujRo2PatGmNzvnb6IqI6N+/fxQXF8fy5csLuVQAgIIrWHht27Ytli9fHlVVVXnjVVVV8eabbzb5eR577LHYtm3bXj/nVV9fH3V1dXk/AAAtTcHCa/PmzRERUV5enjfeoUOH2LRpU5OeY9myZXHVVVfFV77ylejXr1+j+02ePDnKy8tzP927d//wCwcAKJCChVe7du0iIqK2tjZv/L333stdetybt99+O4YPHx7Dhg2LO+64Y6/7TpgwIWpra3M/q1at+vALBwAokILdTqK4uDh69OgRy5YtyxtftmxZ9OrVa69zV6xYEUOHDo2BAwfGww8/HK1bt97r/iUlJVFSUvJ3rxkAoJAK+uH6Cy+8MB599NHcpcWampp44okn4sILL8ztM2fOnLj11ltz2ytXroyhQ4fGKaecEo888sg+7/kFAHCwKOjtJDZs2BBnnnlmtGrVKoYNGxYzZ86MI488Mp599tlo27ZtRETcdNNNcfvtt8d7770XERG9evWK1atXx7hx4/Kia+jQobt947ExbicBAKSyP91R0NNJHTt2jJdeeikee+yxWLlyZUyePHm3O9cPHTo0F2EREZdddlk0NDQUclkAAM2ioGe8moszXgBAKi3mBqoAAPwv4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgETaNPcCAAAKZcfOLBYsXx81G7dGZWnbGNSjU7RuVdRs6xFeAMAhadaS6pj0xNKort2aG+ta3jYmju4TI/t2bZY1udQIABxyZi2pjqsffDkvuiIi1tRujasffDlmLalulnUJLwDgkLJjZxaTnlga2R4e2zU26YmlsWPnnvYoLOEFABxSFixfv9uZrr+VRUR17dZYsHx9ukX9/4QXAHBIqdnYeHR9mP0OJOEFABxSKkvbHtD9DiThBQAcUgb16BRdy9tGYzeNKIoPvt04qEenlMuKCOEFABxiWrcqiomj+0RE7BZfu7Ynju7TLPfzEl4AwCFnZN+uMXXsydGlPP9yYpfytjF17MnNdh8vN1AFAA5JI/t2jRF9urhzPQBACq1bFcXpPSuaexk5LjUCACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkEjBw2v+/Pnx2c9+NgYMGBCXXHJJLFmypCBzAABauoKG1yuvvBLDhg2L448/Pm6//fZo3759nHHGGfH2228f0DkAAAeDoizLskI9+YUXXhjr16+P2bNnR0RElmVxwgknxIgRI+LOO+88YHP+r7q6uigvL4/a2tooKys7MG8GAGAP9qc7CnrGa/bs2TFq1KjcdlFRUYwaNSoXVQdqDgDAwaBNoZ548+bNsWHDhujatWveeNeuXWPVqlUHbE5ERH19fdTX1+e26+rq/o6VAwAURsHOeDU0NERERHFxcd54SUlJbN++/YDNiYiYPHlylJeX5366d+/+9ywdAKAgChZepaWlUVxcHOvWrcsbX7duXXTu3PmAzYmImDBhQtTW1uZ+9nZ2DACguRQsvFq1ahUf+9jH4sUXX8wbf/7552PAgAEHbE7EB2fEysrK8n4AAFqagn64/qqrropp06bFwoULIyLi6aefjjlz5sSVV16Z22fKlCkxePDg/ZoDAHAwKtiH6yMirrjiinjzzTdjyJAhUVFRERs2bIjJkyfnfWuxpqYmli5dul9zAAAORgW9j9cuGzdujDVr1sRRRx0VRxxxRN5jNTU1sW7duujdu3eT5+yL+3gBAKnsT3cU9IzXLqWlpVFaWrrHxyorK6OysnK/5gAAHIz8I9kAAIkILwCARIQXAEAiwgsAIBHhBQCQiPACAEhEeAEAJCK8AAASEV4AAIkILwCARIQXAEAiwgsAIBHhBQCQiPACAEhEeAEAJNKmuRdwsNmxM4sFy9dHzcatUVnaNgb16BStWxU197IAgIOA8NoPs5ZUx6QnlkZ17dbcWNfytjFxdJ8Y2bdrM64MADgYuNTYRLOWVMfVD76cF10REWtqt8bVD74cs5ZUN9PKAICDhfBqgh07s5j0xNLI9vDYrrFJTyyNHTv3tAcAwAeEVxMsWL5+tzNdfyuLiOrarbFg+fp0iwIADjrCqwlqNjYeXR9mPwDgH5PwaoLK0rYHdD8A4B+T8GqCQT06RdfyttHYTSOK4oNvNw7q0SnlsgCAg4zwaoLWrYpi4ug+ERG7xdeu7Ymj+7ifFwCwV8KriUb27RpTx54cXcrzLyd2KW8bU8ee7D5eAMA+uYHqfhjZt2uM6NPFnesBgA9FeO2n1q2K4vSeFc29DADgIORSIwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEhFeAACJCC8AgESEFwBAIsILACAR4QUAkIjwAgBIRHgBACQivAAAEkkSXps2bYply5bFli1bmjzn3Xffjerq6gKuCgAgrYKH1w033BCdO3eOoUOHRkVFRXz/+9/f6/4PPfRQnHjiidGnT5/o379/dO/ePR577LFCLxMAoOAKGl73339/3HXXXTFv3rz4y1/+EjNmzIgJEybEU0891eicl19+OaZNmxbvvvturFmzJsaNGxeXXHJJvPHGG4VcKgBAwRVlWZYV6slPPfXU6N27d/z3f/93bmz48OHRvn37mDFjRpOeY8eOHXH44YfHlClT4otf/GKT5tTV1UV5eXnU1tZGWVnZh1g5AEDT7E93FOyM186dO+PVV1+NU089NW988ODBsXDhwiY/z1tvvRXbt2+P7t27H+glAgAk1WZ/dq6pqYmampq97tOjR49o165dbNy4MbZt2xYVFRV5j1dUVMTatWub9Hrbtm2LL37xizFgwIAYPnx4o/vV19dHfX19bruurq5Jzw8AkNJ+hdfPfvazuPvuu/e6zwMPPBADBw6MNm0+eOpt27blPV5fXx+HHXbYPl+roaEhLr300njnnXdi7ty50bp160b3nTx5ckyaNKkJ7wAAoPnsV3h99atfja9+9atN2rddu3bRsWPH3W4JUV1dvc/Lhjt27IixY8fGggULYs6cOfHRj350r/tPmDAhrr322tx2XV2dS5MAQItT0G81Dhs2LGbOnJnbzrIsZs6cGcOGDcuN1dTUxOuvv57b3hVdzz//fMyZMyeOO+64fb5OSUlJlJWV5f0AALQ0BQ2vG2+8MebPnx/XX399zJ07N770pS/FX//61/jGN76R22fKlClx+umn57Y///nPxy9/+cu488474/33348lS5bEkiVL9vnZMgCAlm6/LjXur/79+8fs2bPjtttui6997WtRVVUVc+fOjWOPPTa3T2VlZfTp0ye3/dprr8UxxxwT3/zmN/Oe65prrolrrrmmkMsFACiogt7Hq7m4jxcAkEqLuI8XAAD5hBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCJtCv0CW7ZsiccffzxWrFgRVVVVccEFF0SbNk172RUrVsQDDzwQffr0iYsvvrjAKwUAKKyiLMuyQj35hg0b4qyzzoqioqIYNmxYzJw5M4488sh49tlno23btnud29DQEGeddVa8/vrr8YlPfCKmTZvW5Netq6uL8vLyqK2tjbKysr/3bQAANGp/uqOglxq/973vxebNm+P555+PO+64I+bNmxevvfZaTJ06dZ9zb7zxxjjuuONiyJAhhVwiAEAyBQ2v6dOnx8UXXxzt27ePiIjKysoYPXr0Ps9ePfvss/Hoo4/GXXfdVcjlAQAkVbDPeG3bti2WL18eVVVVeeNVVVUxc+bMRufV1NTE5ZdfHg8//HCUl5c36bXq6+ujvr4+t11XV/fhFg0AUED7FV6/+c1v4rnnntvrPpdffnkce+yxsXnz5oiI3eKpQ4cOsWnTpj3OzbIsLrvssrj88svjzDPPbPK6Jk+eHJMmTWry/gAAzaFglxrbtWsXERG1tbV54++9917u0uP/9fTTT8ecOXMiy7K46aab4qabboq33norli5dGjfddFO89957e5w3YcKEqK2tzf2sWrXqgL4XAIADYb/OeJ199tlx9tlnN2nf4uLi6NGjRyxbtixvfNmyZdGrV689zjnmmGNiwoQJ+7OkiIgoKSmJkpKS/Z4HAJBSQW8ncd1118X06dPjD3/4Q7Rv3z5qamqiV69e8e1vfzuuvfbaiIiYM2dOvPDCC/H//t//2+NznHfeedG2bVu3kwAAWqQWczuJb37zm3HEEUfE4MGDY9y4cTFkyJA48cQT45prrsntM2fOnLj11lsLuQwAgBahoHeu79ixY7z00kvx2GOPxcqVK2Py5Mm73bl+6NChe72Z6qWXXtrkO90DALRkBb3U2FxcagQAUmkxlxoBAPhfwgsAIBHhBQCQiPACAEhEeAEAJCK8AAASEV4AAIkILwCARIQXAEAiwgsAIBHhBQCQiPACAEhEeAEAJCK8AAASEV4AAIkILwCARIQXAEAiwgsAIBHhBQCQiPACAEhEeAEAJCK8AAASEV4AAIkILwCARIQXAEAiwgsAIBHhBQCQiPACAEhEeAEAJCK8AAASEV4AAIkILwCARIQXAEAiwgsAIBHhBQCQiPACAEhEeAEAJCK8AAASEV4AAIkILwCARIQXAEAiwgsAIBHhBQCQiPACAEhEeAEAJCK8AAASEV4AAIkILwCARIQXAEAiwgsAIBHhBQCQiPACAEhEeAEAJCK8AAASEV4AAIkILwCARIQXAEAiwgsAIJGCh9f8+fPjs5/9bAwYMCAuueSSWLJkyT7nbNmyJW677bYYOnRoDBs2LO65555CLxMAoOAKGl6vvPJKDBs2LI4//vi4/fbbo3379nHGGWfE22+/3eic+vr6+MQnPhEPPfRQfOMb34hbb701Xn311fj5z39eyKUCABRcUZZlWaGe/MILL4z169fH7NmzIyIiy7I44YQTYsSIEXHnnXfucc5tt90Wt9xySyxbtiwqKytz49u3b4/DDjusSa9bV1cX5eXlUVtbG2VlZX//GwEAaMT+dEdBz3jNnj07Ro0aldsuKiqKUaNG5UJsTx566KG44IIL8qIrIpocXQAALVXBwmvz5s2xYcOG6Nq1a954165dY9WqVY3Oe+ONN6J3795x/fXXx8CBA+Pcc8+NBx54IPZ2Yq6+vj7q6uryfgAAWpo2+7Pzj370o7j77rv3us8DDzwQAwcOjIaGhoiIKC4uznu8pKQktm/fvse5WZbFtm3b4uabb45rr7027rzzzli6dGmMHz8+VqxYETfddNMe502ePDkmTZq0P28FACC5/QqvMWPGxLBhw/a6T48ePSIiorS0NIqLi2PdunV5j69bty46d+68x7lFRUVRUVER/fr1i//8z/+MiIhTTz01Vq5cGVOmTGk0vCZMmBDXXnttbruuri66d+/e1LcFAJDEfoVXZWXlbp+9akyrVq3iYx/7WLz44otx9dVX58aff/75GDBgQKPzBg4cGO3atcsb69y5c2zatCmyLIuioqLd5pSUlERJSUkT3wUAQPMo6Ifrr7rqqpg2bVosXLgwIiKefvrpmDNnTlx55ZW5faZMmRKDBw/ObV9zzTXxq1/9KhYvXhwREevXr4/77rsvRo4cucfoAgA4WOzXGa/9dcUVV8Sbb74ZQ4YMiYqKitiwYUNMnjw575uONTU1sXTp0tz2eeedFzfddFOcccYZUVFREWvWrInhw4fHj3/840IuFQCg4Ap6H69dNm7cGGvWrImjjjoqjjjiiLzHampqYt26ddG7d++88a1bt8Y777wTXbp0ifbt2+/X67mPFwCQyv50R0HPeO1SWloapaWle3yssc+NtW3bNo4//vhCLw0AIBn/SDYAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBE2jT3AgD+0ezYmcWC5eujZuPWqCxtG4N6dIrWrYqae1lAAsILIKFZS6pj0hNLo7p2a26sa3nbmDi6T4zs27UZVwak4FIjQCKzllTH1Q++nBddERFrarfG1Q++HLOWVDfTyoBUhBdAAjt2ZjHpiaWR7eGxXWOTnlgaO3buaQ/gUCG8ABJYsHz9bme6/lYWEdW1W2PB8vXpFgUkJ7wAEqjZ2Hh0fZj9gIOT8AJIoLK07QHdDzg4CS+ABAb16BRdy9tGYzeNKIoPvt04qEenlMsCEhNeAAm0blUUE0f3iYjYLb52bU8c3cf9vOAQJ7wAEhnZt2tMHXtydCnPv5zYpbxtTB17svt4wT8AN1AFSGhk364xok8Xd66Hf1DCCyCx1q2K4vSeFc29DKAZuNQIAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgkTbNvYBCyLIsIiLq6uqaeSUAwKFuV2/s6o+9OSTDa+PGjRER0b1792ZeCQDwj2Ljxo1RXl6+132Ksqbk2UFm586dsXr16igtLY2ioqLmXk5B1NXVRffu3WPVqlVRVlbW3MtpURybvXN8GufYNM6x2TvHp3H/CMcmy7LYuHFjdOvWLVq12vunuA7JM16tWrWKo48+urmXkURZWdkh+4v893Js9s7xaZxj0zjHZu8cn8Yd6sdmX2e6dvHhegCARIQXAEAiwusgVVJSEhMnToySkpLmXkqL49jsnePTOMemcY7N3jk+jXNs8h2SH64HAGiJnPECAEhEeAEAJCK8AAASOSTv43WoW7RoUfz5z3+O4447Lk466aQmzdmyZUvMnz8/duzYEYMHD4527doVeJXNY/Xq1fHSSy9FeXl5DB48OIqLi5s8d+bMmVFXVxdjxow5JG+8u2XLlvjd734XW7dujdNPPz06d+68zzkrV66MxYsXR6dOneLkk08+6D8cm2VZLFy4MFatWhUnnHBC9O7duyBzDlYrV66MV155JTp16hSnn356tGmz9/+JaGhoiEWLFkV1dXVUVVVFr169Eq00vU2bNsW8efOioaEhhgwZEh06dGjy3N///vfxxz/+MYYNGxZHHnlk4RbZTHbu3BkLFiyINWvWxIknnhhVVVVNmldbWxvz58+Pww8/PE4//fT9+u/rg1rGQWP79u3ZRRddlHXs2DE755xzso4dO2YXXXRRtn379r3O+8UvfpFVVFRkAwYMyEaPHp316dMne/nllxOtOp2pU6dmRxxxRDZ06NDs+OOPz44//vhs+fLlTZr7s5/9LCsuLs4iYp/H82C0aNGirFu3btmJJ56YDRkyJGvXrl326KOPNrr/6tWrs1GjRmXHHntsdu6552a9e/fOPvrRj2bz589PuOoD6/33389GjBiRVVZWZuecc05WWlqafelLX8p27tx5QOccrP7rv/4rO/zww7Ozzz47O+aYY7K+fftm1dXVje4/Y8aMrKqqKjv55JOzc889N+vQoUN2/vnnZ1u2bEm46jTmz5+fde7cOfvYxz6WnXbaaVlZWVk2c+bMJs1dtmxZ1rlz5ywismeeeabAK02vtrY2Gzx4cNatW7dsxIgRWbt27bKvf/3r+5w3derUrH379tmZZ56ZfepTn8oGDBiQrVy5MsGKm5/wOoj86Ec/yjp06JCLiT/96U9ZWVlZdueddzY6Z/HixVlxcXF2xx135MZWr16dLViwoNDLTeqtt97K2rRpk/30pz/NsizLtm3blp155pnZyJEj9zn3z3/+c3bUUUdlEydOPGTD66STTsouuuiiXDBMnjw5a9++ffbuu+/ucf8333wze/LJJ3PbO3fuzMaOHZv16NEjyXoLYeLEiVm3bt2yNWvWZFn2QYwWFxdnjzzyyAGdczBauHBhVlRUlD3xxBNZln0QnP37988uueSSRuc8/vjj2dtvv53b/stf/pJ17tw5mzRpUsHXm1JDQ0N23HHHZV/4whdyYzfccEPWuXPnbNOmTXudW19fnw0YMCC77bbbDtnwGjduXNazZ89sw4YNWZZl2fPPP58VFRVls2bNanTOk08+mff7lmVZ9vrrr2dvvvlmoZfbIgivg8ipp56aff7zn88b+/d///fstNNOa3TO5z73uaxPnz6FXlqz++53v5tVVlZmO3bsyI09+uijWVFRUVZTU9PovG3btmWnnnpqdu+992Y//elPD8nwWrx4cRYR2bx583JjdXV1WUlJSXbvvfc2+XmmT5+eRURWW1tbiGUWXM+ePbPrrrsub2zUqFHZeeedd0DnHIy+/vWvZ8cff3ze2E9+8pOsuLg427x5c5OfZ/To0dlnPvOZA728ZjV37twsIrLXXnstN1ZdXZ21atUqmzZt2l7njh8/Phs7dmy2atWqQza8OnfunN1yyy15Y4MHD87Gjh3b6JwhQ4Zk559/fqGX1mL5cP1BZPHixdG3b9+8sX79+sXixYsbnfPcc8/FOeecEzU1NfE///M/MW/evNiyZUuhl5rc4sWL48QTT8z7x0n79esXWZbFa6+91ui8b33rW9GtW7f4whe+kGKZzWLX78ff/u6UlpbGscceu9ffnf/r2Wefje7dux+U/9ba+++/H3/605/26+/nw8w5WC1evDj69euXN9avX7/Ytm1bvPXWW016jk2bNsWLL7642/E62C1evDhat26d99m+Ll26xEc+8pG9/h7MnDkzZsyYEXfddVeKZTaL6urqWLt27X79jWzbti1efPHFOOecc2LFihXxi1/8IhYsWBA7duxIseQWwYfrm9Hrr78eixYt2us+Z511VnTr1i0aGhri/fffj06dOuU9XlFREZs3b46GhoY9fhC2pqYm3nzzzRg0aFD069cv/vSnP8WmTZti+vTpMXDgwAP6fg6k2traeOqpp/a6z/HHHx+nnHJKbv89HZuIiPfee2+P859++ul46KGH9vmfQUv05JNPxsaNGxt9/PDDD49Pf/rTEfHBsWnduvVuwVRRUdHosfm/Zs6cGXfffXc8+OCDH3rNzamuri4iYo+/I40dgw8z52BVW1sb3bt3zxvb19/P38qyLK688spo06ZN/Md//Echlthsamtro0OHDrt94WZvvwfV1dXxhS98IX7+859HWVlZ7nfpUFNbWxsR+/c3sn79+mhoaIhf//rX8YMf/CD69esXixYtitLS0vjlL38Zxx57bIFX3fyEVzNatmxZzJgxY6/79OrVK7p16xZt2rSJ1q1bx6ZNm/Ie37RpU7Rp06bRbx+1bds2XnzxxViyZEl07do1du7cGRdffHFcccUVLfr/tW/cuHGfx2bEiBG58CopKdnjsYn44BjsyRVXXBHnnXdePPPMMxER8cILL0RExKOPPhr9+/dv0d9ee+aZZ2LNmjWNPt6hQ4dceJWUlMSOHTti69atecdi06ZNjR6bvzV37tz4l3/5l7jpppviX//1X//+xTeDXd/G3NPvSGPH4MPMOVh9mL+fvzVu3Lj41a9+Fb/5zW+a9G3Zg0lJSUls3rx5t/G9/R5cd911UVVVFe+880488sgjsX79+oiImDNnThx22GHx8Y9/vKBrTuXD/I3sGn/99ddjyZIl0a5du6ivr48zzjgjxo8fH48//nhhF90CCK9mdP7558f555/f5P2PO+64WLlyZd7YihUrokePHo3O6dmzZ1RWVkbXrl0jIqJVq1bxmc98Ji677LLd/oe4JTn66KPjkUceafL+PXv2jFmzZuWNrVixIiI+OG57cvbZZ8d7772XC7xd+//iF7+I9u3bt+jwuv3225u8b8+ePSPig1sF/NM//VNEfPD173feeafRY7PL7373uxg1alRcd9118e1vf/tDr7e5dezYMTp27LjHv5/GjsGHmXOw6tmzZ7zxxht5Y/v6+9ll/Pjx8eCDD8azzz7b5NvbHEx69uwZW7dujZqamqisrIyIiK1bt8Zf//rXRo9N3759o6GhIfffLbs+3jFv3rxo3779IRNeRx99dBQXF+/X30iHDh2ioqIiRo4cmbutUUlJSZx33nlx3333FXzNLUJzf8iMphs3blxWVVWVbdu2LcuyD74x07Nnz+xrX/tabp833ngj+9nPfpbb/u53v5v16tUr70Pn3/rWt7LKysp0C0/gmWeeySIiW7x4cW7sy1/+ct4Hhuvq6rKHH3640a/IH6ofrq+vr886deqU3XzzzbmxWbNm7Xa8nnzyyeyll17Kbc+bNy9r37599p3vfCfpegvl3/7t37JBgwbl/hY2bdqUde7cOe+4LFq0KJsxY8Z+zTkU/PznP89atWqVrVixIjd2ySWXZKecckpue+3atdnDDz+crV27Njc2fvz4rGPHjtnvf//7pOtNqba2NjviiCPyvj3+yCOPZK1bt847XjNmzMgWLVq0x+c4lD9cf+6552bDhw/Pba9bty5r165d3vF66aWX8r4l/bnPfW63b5yPGTMmGzx4cOEX3AIIr4NIdXV1dtRRR2XnnHNONmXKlGzEiBHZUUcdlRcS3//+97PWrVvntuvq6rLevXtn5557bnbvvfdmN9xwQ9a2bdvsvvvua463UFCf/vSnsx49emR33HFHNm7cuKxNmzbZL3/5y9zjr7/+ehYR2VNPPbXH+YdqeGVZlt1///3ZYYcdln3rW9/KfvCDH2SVlZXZVVddlbdPr169cmOvv/56Vlpamg0aNCh7+OGH837q6uqa4y383f74xz9mnTp1yj7zmc9kU6ZMyYYMGZJVVVXlfUvzhhtuyI488sj9mnMo2LFjRzZs2LCsd+/e2Y9+9KPsyiuvzIqLi7Pf/va3uX3mz5+fRUTuXm633HJLFhHZddddl/f7cSjGxQ9/+MPs8MMPzyZNmpTdeuutWYcOHbIbbrghb58jjzxyt7FdDuXw+sMf/pC1b98+u/TSS7O77rorGzBgQHbSSSfl3c/tqquuynr16pXbfvvtt7Mjjzwyu/zyy7P77rsv+/KXv5wVFxdnv/71r5vjLSTnUuNBpEuXLrFw4cKYOnVqzJ8/PwYPHhw//elP8+6EfMIJJ8SYMWNy26WlpfHCCy/E3XffHXPnzo0jjzwynnvuuRb9wfoPa9q0aXH//ffH/Pnzo6ysLJ5//vm891lWVhZjxozJXXb9v4499tgYM2ZM3jcjDxWf//zno0ePHvHoo49GdXV1/PCHP4xLL700b59zzz03+vTpExEffEZj1KhRERG7fdburLPOitLS0iTrPpB69uwZr7zyStx9993xwgsvxKhRo+Kaa67J+9LBSSedFBdccMF+zTkUtGrVKp566qm455574ve//3106tQpXnrppfjnf/7n3D6dO3eOMWPG5D7DVVJSEmPGjImVK1fmXWqqqqqK4cOHJ38PhTR+/Pjo3bt3zJgxIxoaGuKee+6Jiy66KG+fCy64oNFLrUcccUSMGTMmunTpkmK5SfXr1y9eeeWVuOeee+LFF1+MMWPGxNVXX533MZaBAwfm/WspxxxzTO7v6rnnnovu3bvHkiVLmnzH+4NdUZZlWXMvAgDgH8Gh93/tAQBaKOEFAJCI8AIASER4AQAkIrwAABIRXgAAiQgvAIBEhBcAQCLCCwAgEeEFAJCI8AIASER4AQAk8v8Bg9cMuPftgcAAAAAASUVORK5CYII=",
      "text/plain": [
       "<Figure size 700x700 with 1 Axes>"
      ]
     },
     "metadata": {},
     "output_type": "display_data"
    }
   ],
   "source": [
    "fig, ax = plt.subplots(figsize=(7, 7))\n",
    "ax.scatter(x[:,0], x[:,1])"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "c10ac9de",
   "metadata": {},
   "source": [
    "## Reconstructing the environment"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 15,
   "id": "357a22b0",
   "metadata": {},
   "outputs": [],
   "source": [
    "exact_V = -np.array([[0,1,1,np.sqrt(2)],[1,0,np.sqrt(2),1], [1,np.sqrt(2),0,1], [np.sqrt(2),1,1,0]])\n",
    "Pa = np.array([P[:,:,a] for a in range(n_actions)])\n",
    "Ra = np.array([-np.ones((n_states, n_states)) for a in range(n_actions)])\n",
    "mask = np.ones((n_states, n_states)) - np.eye(n_states) # implements boundary condition\n",
    "Qa = np.array([(Ra[a] +  Pa[a] @ V)*mask  for a in range(n_actions)])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1ce3aa5e",
   "metadata": {},
   "outputs": [],
   "source": [
    "Pa_reconstructed = []\n",
    "for a in range(n_actions):\n",
    "    Pa_recon = []\n",
    "    for s in range(n_states):\n",
    "        V_s = np.delete(exact_V, s, axis=0)\n",
    "        q_ss = np.delete(Qa[a,:,s], s, axis=0) \n",
    "        r_ss = np.delete(Ra[a,:,s], s, axis=0)\n",
    "        Pa_recon.append(\\\n",
    "            np.linalg.pinv(np.vstack([V_s, np.ones(n_states)])) @ \\\n",
    "                np.concatenate([q_ss - r_ss, [1]])\n",
    "        )\n",
    "    Pa_reconstructed.append(Pa_recon)\n",
    "Pa_reconstructed = np.array(Pa_reconstructed)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 17,
   "id": "3126380a",
   "metadata": {},
   "outputs": [
    {
     "data": {
      "text/plain": [
       "array([[[ 0.   ,  1.   , -0.   ,  0.   ],\n",
       "        [ 1.   , -0.   , -0.   ,  0.   ],\n",
       "        [-0.   ,  0.   ,  0.   ,  1.   ],\n",
       "        [ 0.   , -0.   ,  1.   ,  0.   ]],\n",
       "\n",
       "       [[-0.   ,  0.   ,  1.   , -0.   ],\n",
       "        [ 0.   , -0.   , -0.   ,  1.   ],\n",
       "        [ 1.   , -0.   , -0.   ,  0.   ],\n",
       "        [-0.   ,  1.   ,  0.   , -0.   ]],\n",
       "\n",
       "       [[ 0.293, -0.   , -0.   ,  0.707],\n",
       "        [ 0.   ,  0.293,  0.707,  0.   ],\n",
       "        [ 0.   ,  0.707,  0.293,  0.   ],\n",
       "        [ 0.707,  0.   , -0.   ,  0.293]]])"
      ]
     },
     "execution_count": 17,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "Pa_reconstructed"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "0f97b128",
   "metadata": {},
   "source": [
    "Doing it simultaneously for all actions is more efficient."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 21,
   "id": "f234bbdd",
   "metadata": {},
   "outputs": [
    {
     "data": {
      "text/plain": [
       "array([[[ 0.   ,  1.   , -0.   ,  0.   ],\n",
       "        [ 1.   , -0.   ,  0.   , -0.   ],\n",
       "        [-0.   , -0.   , -0.   ,  1.   ],\n",
       "        [ 0.   ,  0.   ,  1.   ,  0.   ]],\n",
       "\n",
       "       [[-0.   ,  0.   ,  1.   , -0.   ],\n",
       "        [ 0.   , -0.   , -0.   ,  1.   ],\n",
       "        [ 1.   , -0.   , -0.   ,  0.   ],\n",
       "        [-0.   ,  1.   ,  0.   , -0.   ]],\n",
       "\n",
       "       [[ 0.293,  0.   ,  0.   ,  0.707],\n",
       "        [ 0.   ,  0.293,  0.707,  0.   ],\n",
       "        [-0.   ,  0.707,  0.293, -0.   ],\n",
       "        [ 0.707,  0.   ,  0.   ,  0.293]]])"
      ]
     },
     "execution_count": 21,
     "metadata": {},
     "output_type": "execute_result"
    }
   ],
   "source": [
    "Pa_reconstructed = []\n",
    "for s in range(n_states):\n",
    "    V_s = np.delete(exact_V, s, axis=0)\n",
    "    q_ss = np.delete(Qa[:,:,s], s, axis=1) \n",
    "    r_ss = np.delete(Ra[:,:,s], s, axis=1)\n",
    "    Pa_reconstructed.append(\\\n",
    "        np.linalg.pinv(np.vstack([V_s, np.ones(n_states)])) @ \\\n",
    "            np.vstack([(q_ss - r_ss).T, np.ones(n_states-1)]))\n",
    "Pa_reconstructed = np.transpose(np.array(Pa_reconstructed), (2,1,0))\n",
    "Pa_reconstructed"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "qbist_spacetime",
   "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.14.6"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
