{ "cells": [ { "cell_type": "markdown", "id": "0af7e8f5", "metadata": {}, "source": [ "# Museum visitors vs. city population\n", "\n", "This notebook talks to the `museum_api` FastAPI server over HTTP: it starts the server (or uses one you already run), sends it the data, and plots the regression it returns.\n", "\n", "**Setup:** put `../api/app/museum_api.py` next to this notebook, then `pip install scikit-learn pandas fastapi uvicorn requests matplotlib`." ] }, { "cell_type": "code", "id": "4c50f329", "metadata": { "ExecuteTime": { "end_time": "2026-10-07T21:11:58.303386100Z", "start_time": "2026-10-07T21:11:57.570566Z" } }, "source": [ "import threading, time\n", "import requests\n", "import numpy as np\n", "import pandas as pd\n", "import matplotlib.pyplot as plt\n", "from matplotlib.ticker import FuncFormatter\n", "\n", "API = \"http://127.0.0.1:8000\" # change if your server runs elsewhere" ], "outputs": [ { "ename": "ModuleNotFoundError", "evalue": "No module named 'matplotlib'", "output_type": "error", "traceback": [ "\u001B[31m---------------------------------------------------------------------------\u001B[39m", "\u001B[31mModuleNotFoundError\u001B[39m Traceback (most recent call last)", "\u001B[36mCell\u001B[39m\u001B[36m \u001B[39m\u001B[32mIn[1]\u001B[39m\u001B[32m, line 5\u001B[39m\n\u001B[32m 1\u001B[39m \u001B[38;5;28;01mimport\u001B[39;00m threading, time\n\u001B[32m 2\u001B[39m \u001B[38;5;28;01mimport\u001B[39;00m requests\n\u001B[32m 3\u001B[39m \u001B[38;5;28;01mimport\u001B[39;00m numpy \u001B[38;5;28;01mas\u001B[39;00m np\n\u001B[32m 4\u001B[39m \u001B[38;5;28;01mimport\u001B[39;00m pandas \u001B[38;5;28;01mas\u001B[39;00m pd\n\u001B[32m----> \u001B[39m\u001B[32m5\u001B[39m \u001B[38;5;28;01mimport\u001B[39;00m matplotlib.pyplot \u001B[38;5;28;01mas\u001B[39;00m plt\n\u001B[32m 6\u001B[39m \u001B[38;5;28;01mfrom\u001B[39;00m matplotlib.ticker \u001B[38;5;28;01mimport\u001B[39;00m FuncFormatter\n\u001B[32m 7\u001B[39m \n\u001B[32m 8\u001B[39m API = \u001B[33m\"http://127.0.0.1:8000\"\u001B[39m \u001B[38;5;66;03m# change if your server runs elsewhere\u001B[39;00m\n", "\u001B[31mModuleNotFoundError\u001B[39m: No module named 'matplotlib'" ] } ], "execution_count": 1 }, { "cell_type": "markdown", "id": "69524967", "metadata": {}, "source": [ "## 1. Start the server\n", "\n", "If a server is already answering at `API` (e.g. you started it with `uvicorn museum_api:app` in a terminal), this cell just uses it. Otherwise it launches one in a background thread inside this kernel. Restart the kernel to stop it." ] }, { "cell_type": "code", "execution_count": 2, "id": "b9d41502", "metadata": {}, "outputs": [ { "ename": "NameError", "evalue": "name 'API' is not defined", "output_type": "error", "traceback": [ "\u001B[31m---------------------------------------------------------------------------\u001B[39m", "\u001B[31mNameError\u001B[39m Traceback (most recent call last)", "\u001B[36mCell\u001B[39m\u001B[36m \u001B[39m\u001B[32mIn[2]\u001B[39m\u001B[32m, line 7\u001B[39m\n\u001B[32m 3\u001B[39m \u001B[38;5;28;01mreturn\u001B[39;00m requests.get(f\"{API}/health\", timeout=\u001B[32m1\u001B[39m).ok\n\u001B[32m 4\u001B[39m \u001B[38;5;28;01mexcept\u001B[39;00m requests.ConnectionError:\n\u001B[32m 5\u001B[39m \u001B[38;5;28;01mreturn\u001B[39;00m \u001B[38;5;28;01mFalse\u001B[39;00m\n\u001B[32m 6\u001B[39m \n\u001B[32m----> \u001B[39m\u001B[32m7\u001B[39m \u001B[38;5;28;01mif\u001B[39;00m \u001B[38;5;28;01mnot\u001B[39;00m server_up():\n\u001B[32m 8\u001B[39m \u001B[38;5;28;01mimport\u001B[39;00m uvicorn\n\u001B[32m 9\u001B[39m \u001B[38;5;28;01mfrom\u001B[39;00m museum_api \u001B[38;5;28;01mimport\u001B[39;00m create_app\n\u001B[32m 10\u001B[39m server = uvicorn.Server(uvicorn.Config(create_app(), host=\u001B[33m\"127.0.0.1\"\u001B[39m, port=\u001B[32m8000\u001B[39m, log_level=\u001B[33m\"warning\"\u001B[39m))\n", "\u001B[36mCell\u001B[39m\u001B[36m \u001B[39m\u001B[32mIn[2]\u001B[39m\u001B[32m, line 4\u001B[39m, in \u001B[36mserver_up\u001B[39m\u001B[34m()\u001B[39m\n\u001B[32m 1\u001B[39m \u001B[38;5;28;01mdef\u001B[39;00m server_up():\n\u001B[32m 2\u001B[39m \u001B[38;5;28;01mtry\u001B[39;00m:\n\u001B[32m 3\u001B[39m \u001B[38;5;28;01mreturn\u001B[39;00m requests.get(f\"{API}/health\", timeout=\u001B[32m1\u001B[39m).ok\n\u001B[32m----> \u001B[39m\u001B[32m4\u001B[39m \u001B[38;5;28;01mexcept\u001B[39;00m requests.ConnectionError:\n\u001B[32m 5\u001B[39m \u001B[38;5;28;01mreturn\u001B[39;00m \u001B[38;5;28;01mFalse\u001B[39;00m\n", "\u001B[31mNameError\u001B[39m: name 'API' is not defined" ] } ], "source": [ "def server_up():\n", " try:\n", " return requests.get(f\"{API}/health\", timeout=1).ok\n", " except requests.ConnectionError:\n", " return False\n", "\n", "if not server_up():\n", " import uvicorn\n", " from api.app.museum_api import create_app\n", " server = uvicorn.Server(uvicorn.Config(create_app(), host=\"127.0.0.1\", port=8000, log_level=\"warning\"))\n", " threading.Thread(target=server.run, daemon=True).start()\n", " for _ in range(50):\n", " if server_up(): break\n", " time.sleep(0.1)\n", "\n", "requests.get(f\"{API}/health\").json()" ] }, { "cell_type": "markdown", "id": "742501f7", "metadata": {}, "source": [ "## 2. Load your data\n", "\n", "Replace this with however you build your DataFrame. It needs the columns `Name, Visitors, City, Country, Population`." ] }, { "cell_type": "code", "execution_count": null, "id": "cb58b0f5", "metadata": {}, "outputs": [], "source": [ "df = pd.read_csv(\"museums.csv\")\n", "df.head()" ] }, { "cell_type": "markdown", "id": "362d6127", "metadata": {}, "source": [ "## 3. Train the model through the API\n", "\n", "`scale=\"log\"` fits log(visitors) ~ log(population). Set `aggregate_by_city=True` to count each city once instead of once per museum." ] }, { "cell_type": "code", "execution_count": null, "id": "ab229859", "metadata": {}, "outputs": [], "source": [ "def train(df, scale=\"log\", aggregate_by_city=False):\n", " records = df[[\"Name\", \"Visitors\", \"City\", \"Country\", \"Population\"]].dropna().to_dict(\"records\")\n", " r = requests.post(f\"{API}/train\", json={\"records\": records, \"scale\": scale,\n", " \"aggregate_by_city\": aggregate_by_city})\n", " r.raise_for_status()\n", " return r.json()\n", "\n", "summary = train(df, scale=\"log\")\n", "pd.Series(summary).to_frame(\"value\")" ] }, { "cell_type": "markdown", "id": "f655bbcf", "metadata": {}, "source": [ "## 4. Plot the regression" ] }, { "cell_type": "code", "execution_count": null, "id": "1a23269f", "metadata": {}, "outputs": [], "source": [ "BLUE, ORANGE, INK, MUTED = \"#2a78d6\", \"#eb6834\", \"#0b0b0b\", \"#52514e\"\n", "plt.rcParams.update({\"axes.spines.top\": False, \"axes.spines.right\": False,\n", " \"axes.edgecolor\": MUTED, \"axes.labelcolor\": INK,\n", " \"xtick.color\": MUTED, \"ytick.color\": MUTED,\n", " \"axes.grid\": True, \"axes.axisbelow\": True, \"grid.color\": \"#e6e5e0\", \"grid.linewidth\": 0.8,\n", " \"figure.dpi\": 110})\n", "\n", "def human(x, _=None):\n", " for div, suf in ((1e9, \"B\"), (1e6, \"M\"), (1e3, \"K\")):\n", " if abs(x) >= div: return f\"{x/div:g}{suf}\"\n", " return f\"{x:g}\"\n", "\n", "points = pd.DataFrame(requests.get(f\"{API}/data\").json())\n", "\n", "# Fitted curve: ask the API to predict across the population range\n", "grid = np.geomspace(points.Population.min(), points.Population.max(), 100)\n", "curve = pd.DataFrame(requests.post(f\"{API}/predict\", json={\"populations\": grid.tolist()}).json())\n", "\n", "fig, ax = plt.subplots(figsize=(9, 6))\n", "ax.scatter(points.Population, points.Visitors, s=40, color=BLUE, alpha=0.75,\n", " edgecolor=\"white\", linewidth=1, label=\"Museums\", zorder=3)\n", "ax.plot(curve.population, curve.predicted_visitors, color=ORANGE, lw=2,\n", " label=f\"Fit: {summary['equation']}\", zorder=4)\n", "\n", "# Label the 5 biggest outliers (furthest from the line on the log scale)\n", "points[\"logres\"] = np.log10(points.Visitors / points.Predicted)\n", "for _, r in points.reindex(points.logres.abs().sort_values(ascending=False).index).head(5).iterrows():\n", " ax.annotate(r.Name, (r.Population, r.Visitors), xytext=(6, 4),\n", " textcoords=\"offset points\", fontsize=8, color=MUTED)\n", "\n", "if summary[\"scale\"] == \"log\":\n", " ax.set_xscale(\"log\"); ax.set_yscale(\"log\")\n", "ax.xaxis.set_major_formatter(FuncFormatter(human))\n", "ax.yaxis.set_major_formatter(FuncFormatter(human))\n", "ax.set_xlabel(\"City population\"); ax.set_ylabel(\"Annual visitors\")\n", "cv = summary[\"cv_r2_5fold\"]\n", "ax.set_title(f\"Visitors vs. city population R² = {summary['r2']:.2f}\"\n", " + (f\" (cross-validated {cv:.2f})\" if cv is not None else \"\"),\n", " loc=\"left\", fontsize=12, color=INK)\n", "ax.legend(frameon=False, loc=\"upper left\")\n", "plt.tight_layout(); plt.show()\n", "\n", "print(summary[\"interpretation\"])" ] }, { "cell_type": "markdown", "id": "a4eafc4a", "metadata": {}, "source": [ "## 5. Which museums beat (or miss) their city's size?\n", "\n", "Ratio = actual ÷ predicted visitors. Above 1× means the museum draws more than its city's size alone would suggest." ] }, { "cell_type": "code", "execution_count": null, "id": "9a823747", "metadata": {}, "outputs": [], "source": [ "N = 8\n", "over = pd.DataFrame(requests.get(f\"{API}/residuals\", params={\"top\": N, \"order\": \"over\"}).json())\n", "under = pd.DataFrame(requests.get(f\"{API}/residuals\", params={\"top\": N, \"order\": \"under\"}).json())\n", "res = pd.concat([over, under.iloc[::-1]]).drop_duplicates(\"Name\").iloc[::-1]\n", "\n", "fig, ax = plt.subplots(figsize=(9, 0.35 * len(res) + 1))\n", "colors = [BLUE if r >= 1 else ORANGE for r in res.Ratio]\n", "ax.barh(res.Name + \" (\" + res.City + \")\", np.log10(res.Ratio), color=colors, height=0.7)\n", "ax.axvline(0, color=MUTED, lw=1)\n", "ticks = np.array([0.1, 0.2, 0.5, 1, 2, 5, 10])\n", "ticks = ticks[(np.log10(ticks) >= np.log10(res.Ratio).min() - 0.1) &\n", " (np.log10(ticks) <= np.log10(res.Ratio).max() + 0.1)]\n", "ax.set_xticks(np.log10(ticks), [f\"{t:g}×\" for t in ticks])\n", "ax.set_xlabel(\"Actual ÷ predicted visitors (log scale)\")\n", "ax.set_title(\"Over-performers (blue) and under-performers (orange)\", loc=\"left\", fontsize=12, color=INK)\n", "ax.grid(axis=\"y\", visible=False)\n", "plt.tight_layout(); plt.show()\n", "\n", "res[[\"Name\", \"City\", \"Country\", \"Population\", \"Visitors\", \"Predicted\", \"Ratio\"]].round(2)" ] }, { "cell_type": "markdown", "id": "ce7a12d7", "metadata": {}, "source": [ "## 6. Compare model variants\n", "\n", "Retrains the server with each setting and collects the fit statistics. The last setting tried stays loaded, so the final call re-trains the default." ] }, { "cell_type": "code", "execution_count": null, "id": "9adca34e", "metadata": {}, "outputs": [], "source": [ "rows = []\n", "for scale in (\"linear\", \"log\"):\n", " for agg in (False, True):\n", " s = train(df, scale=scale, aggregate_by_city=agg)\n", " rows.append({k: s[k] for k in (\"scale\", \"aggregate_by_city\", \"n_samples\",\n", " \"r2\", \"cv_r2_5fold\", \"pearson_r\", \"spearman_rho\")})\n", "train(df, scale=\"log\") # restore the default model\n", "pd.DataFrame(rows).round(3)" ] }, { "cell_type": "markdown", "id": "1c0bcd2f", "metadata": {}, "source": [ "## 7. Predict for any city" ] }, { "cell_type": "code", "execution_count": null, "id": "6ff399ee", "metadata": {}, "outputs": [], "source": [ "cities = {\"Montréal\": 1_780_000, \"Lyon\": 520_000, \"Tokyo\": 14_000_000}\n", "pred = requests.post(f\"{API}/predict\", json={\"populations\": list(cities.values())}).json()\n", "pd.DataFrame(pred, index=cities.keys()).round(0)" ] } ], "metadata": { "kernelspec": { "display_name": "Python 3 (ipykernel)", "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.12.10" } }, "nbformat": 4, "nbformat_minor": 5 }