Files
museum_analytics/notebooks/work/museum_regression.ipynb
T
2026-10-08 15:15:39 -04:00

314 lines
13 KiB
Plaintext
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
{
"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
}