{"nbformat":4,"nbformat_minor":0,"metadata":{"colab":{"provenance":[],"authorship_tag":"ABX9TyOvpy1stEUkNIKzyv0u7M2R"},"kernelspec":{"name":"python3","display_name":"Python 3"},"language_info":{"name":"python"}},"cells":[{"cell_type":"code","execution_count":2,"metadata":{"colab":{"base_uri":"https://localhost:8080/","height":898,"output_embedded_package_id":"1h-guCCBuN8pTEOXAYroKdKRrckKK2X12"},"id":"FO9rfZpYaVtn","executionInfo":{"status":"ok","timestamp":1785052670363,"user_tz":-330,"elapsed":29046,"user":{"displayName":"Raunak Bhattacharyya","userId":"05940311940370151459"}},"outputId":"9b3cb10a-b4c6-4a52-fa61-9032d93d9510"},"outputs":[{"output_type":"display_data","data":{"text/plain":"Output hidden; open in https://colab.research.google.com to view."},"metadata":{}}],"source":["import numpy as np\n","import matplotlib.pyplot as plt\n","from matplotlib.animation import FuncAnimation\n","from scipy.special import erfinv\n","from IPython.display import HTML\n","\n","# =====================================================\n","# Parameters\n","# =====================================================\n","\n","N = 1000\n","batch_size = 20\n","\n","# Uniform samples\n","u = np.random.rand(N)\n","\n","# Gaussian samples via inverse transform\n","g = np.sqrt(2) * erfinv(2*u - 1)\n","\n","# True inverse CDF\n","u_curve = np.linspace(0.001, 0.999, 500)\n","g_curve = np.sqrt(2) * erfinv(2*u_curve - 1)\n","\n","# =====================================================\n","# Figure\n","# =====================================================\n","\n","fig, ax = plt.subplots(figsize=(8,8))\n","\n","def update(frame):\n","\n","    ax.clear()\n","\n","    end = min((frame + 1) * batch_size, N)\n","    previous_end = max(0, end - batch_size)\n","\n","    # -------------------------------------------------\n","    # Inverse CDF (fixed)\n","    # -------------------------------------------------\n","\n","    ax.plot(\n","        u_curve,\n","        g_curve,\n","        color=\"royalblue\",\n","        linewidth=3,\n","        label=r\"$F^{-1}(u)$\"\n","    )\n","\n","    # -------------------------------------------------\n","    # Previous samples (grey dots only)\n","    # -------------------------------------------------\n","\n","    ax.scatter(\n","        u[:previous_end],\n","        np.zeros(previous_end),\n","        s=15,\n","        color=\"lightgray\",\n","        alpha=0.6,\n","        zorder=2\n","    )\n","\n","    ax.scatter(\n","        np.zeros(previous_end),\n","        g[:previous_end],\n","        s=15,\n","        color=\"lightgray\",\n","        alpha=0.6,\n","        zorder=2\n","    )\n","\n","    # -------------------------------------------------\n","    # Current batch\n","    # -------------------------------------------------\n","\n","    for ui, gi in zip(u[previous_end:end], g[previous_end:end]):\n","\n","        # Uniform sample on x-axis\n","        ax.scatter(\n","            ui,\n","            0,\n","            s=45,\n","            color=\"darkorange\",\n","            edgecolors=\"black\",\n","            zorder=5\n","        )\n","\n","        # Gaussian sample on y-axis\n","        ax.scatter(\n","            0,\n","            gi,\n","            s=45,\n","            color=\"darkorange\",\n","            edgecolors=\"black\",\n","            zorder=5\n","        )\n","\n","        # Vertical guide to curve\n","        ax.plot(\n","            [ui, ui],\n","            [0, gi],\n","            \"--\",\n","            color=\"gray\",\n","            linewidth=1\n","        )\n","\n","        # Horizontal guide to y-axis\n","        ax.plot(\n","            [0, ui],\n","            [gi, gi],\n","            \"--\",\n","            color=\"gray\",\n","            linewidth=1\n","        )\n","\n","    # -------------------------------------------------\n","\n","    ax.set_xlim(-0.05, 1.05)\n","    ax.set_ylim(-4.2, 4.2)\n","\n","    ax.set_xlabel(\"Uniform sample $u$\")\n","    ax.set_ylabel(r\"Inverse transformed value $F^{-1}(u)$\")\n","\n","    ax.set_title(f\"Inverse Transform Sampling ({end} samples)\")\n","\n","    ax.grid(alpha=0.3)\n","    ax.legend(loc=\"upper left\")\n","\n","ani = FuncAnimation(\n","    fig,\n","    update,\n","    frames=int(np.ceil(N / batch_size)),\n","    interval=400,\n","    repeat=False\n",")\n","\n","plt.close(fig)\n","HTML(ani.to_jshtml())"]}]}