restructure project with seperate app/ and notebooks/ sub directories
This commit is contained in:
1 parent
2ac6ac005a
commit
d5da086e66
16 files changed
+40
-14
No files matched your search
@@ -0,0 +1,313 @@
|
||||
{
|
||||
"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
|
||||
}
|
||||
Reference in new issue
Block a user