Data setup done with museum data from wikipedia and updated with city population
This commit is contained in:
1 parent
ce7b60a8f1
commit
fc3dcf19ba
9 files changed
+50818
No files matched your search
@@ -0,0 +1,2 @@
|
|||||||
|
.env
|
||||||
|
.pyc
|
||||||
@@ -0,0 +1,17 @@
|
|||||||
|
# Museum Analytics #
|
||||||
|
|
||||||
|
## ⚙️ Configuration Setup
|
||||||
|
|
||||||
|
This project uses environment variables to secure sensitive information. Follow these steps to configure your local environment:
|
||||||
|
|
||||||
|
1. **Duplicate the template file:**
|
||||||
|
Copy the `.env.example` file and rename it to `.env` in the root directory.
|
||||||
|
```bash
|
||||||
|
cp .env.example .env
|
||||||
|
```
|
||||||
|
2. **Update your keys:**
|
||||||
|
Open the newly created `.env` file in your text editor and replace the placeholder values with your actual credentials:
|
||||||
|
```text
|
||||||
|
WIKIPEDIA_USER_AGENT="MuseumAnalytics/0.1 (your_email_address)"
|
||||||
|
```
|
||||||
|
⚠️ **Important:** Never commit your `.env` file to Git. It is already added to `.gitignore` to protect your secrets.
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
WIKIPEDIA_USER_AGENT="MuseumAnalytics/0.1 (your_email_address)" # johndoe@gmail.com
|
||||||
Whitespace-only changes.
@@ -0,0 +1,7 @@
|
|||||||
|
|
||||||
|
|
||||||
|
DEFAULT_MUSEUM_DATA_SOURCE_URL = "https://en.wikipedia.org/wiki/List_of_most_visited_museums"
|
||||||
|
|
||||||
|
# csv file from https://simplemaps.com/data/world-cities
|
||||||
|
RAW_POPULATION_DATA_FILE= "data/worldcities.csv"
|
||||||
|
MUSEUM_DATA_FILE= "data/updated_museum_data.csv"
|
||||||
File diff suppressed because it is too large.
Load diff
@@ -0,0 +1,92 @@
|
|||||||
|
import os
|
||||||
|
import requests
|
||||||
|
import wikipediaapi
|
||||||
|
from bs4 import BeautifulSoup
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
from museum_analytics import constants
|
||||||
|
from io import StringIO
|
||||||
|
import pandas as pd
|
||||||
|
|
||||||
|
load_dotenv()
|
||||||
|
|
||||||
|
WIKIPEDIA_USER_AGENT = os.getenv("WIKIPEDIA_USER_AGENT")
|
||||||
|
SOUP_PER_SECTION: dict[str, BeautifulSoup] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def get_first_data_table_from_wikipedia(page_title: str) -> pd.DataFrame:
|
||||||
|
"""
|
||||||
|
|
||||||
|
:param page_title:
|
||||||
|
:return:
|
||||||
|
"""
|
||||||
|
|
||||||
|
page_url = f"https://en.wikipedia.org/api/rest_v1/page/html/{page_title}"
|
||||||
|
print(f"Fetching data from Wikipedia: {page_url}")
|
||||||
|
with requests.Session() as session:
|
||||||
|
session.headers["User-Agent"] = WIKIPEDIA_USER_AGENT
|
||||||
|
html = session.get(page_url, timeout=30).text
|
||||||
|
|
||||||
|
tables = pd.read_html(StringIO(html), attrs={"class": "wikitable"})
|
||||||
|
# TODO: add a way to specifically fetch a table
|
||||||
|
df = tables[0]
|
||||||
|
|
||||||
|
if df.iloc[-1].isna().all():
|
||||||
|
df = df.iloc[:-1]
|
||||||
|
return df
|
||||||
|
|
||||||
|
|
||||||
|
def _get_city_population_data_frame_from_raw_data(museum_cities_df: pd.DataFrame) -> pd.DataFrame:
|
||||||
|
"""
|
||||||
|
|
||||||
|
:param museum_cities_df:
|
||||||
|
:return:
|
||||||
|
"""
|
||||||
|
raw_df = pd.read_csv(constants.RAW_POPULATION_DATA_FILE)
|
||||||
|
raw_df = raw_df.drop(["city_ascii", "lat", "lng", "iso2", "iso3", "capital", "id"], axis=1)
|
||||||
|
raw_df["population"] = raw_df["population"].astype("Int64")
|
||||||
|
wrong_washington_condition = (
|
||||||
|
(raw_df["city"] == "Washington")
|
||||||
|
& (raw_df["country"] == "United States")
|
||||||
|
& (raw_df["admin_name"] != "District of Columbia")
|
||||||
|
)
|
||||||
|
cities_population = raw_df[~wrong_washington_condition]
|
||||||
|
cities_population.rename(columns={"population": "Population"}, inplace=True)
|
||||||
|
return cities_population
|
||||||
|
|
||||||
|
|
||||||
|
def _update_city_population(museum_df) -> pd.DataFrame:
|
||||||
|
city_pop = _get_city_population_data_frame_from_raw_data(museum_df)
|
||||||
|
museum_df = pd.merge(
|
||||||
|
museum_df,
|
||||||
|
city_pop,
|
||||||
|
left_on=["City_Clean", "Country"],
|
||||||
|
right_on=["city", "country"],
|
||||||
|
how="left"
|
||||||
|
)
|
||||||
|
museum_df.drop(columns=["city", "country"], inplace=True)
|
||||||
|
|
||||||
|
return museum_df
|
||||||
|
|
||||||
|
|
||||||
|
def get_museum_data() -> pd.DataFrame:
|
||||||
|
url = constants.DEFAULT_MUSEUM_DATA_SOURCE_URL
|
||||||
|
page = url.rsplit("/", 1)[-1]
|
||||||
|
museum_df = get_first_data_table_from_wikipedia(page_title=page)
|
||||||
|
# print(museum_df.to_string())
|
||||||
|
|
||||||
|
city_corrections = {
|
||||||
|
'Washington, D.C.': 'Washington',
|
||||||
|
'New York City': 'New York',
|
||||||
|
'Vatican City, Rome': 'Vatican City',
|
||||||
|
'London, South Kensington': 'London',
|
||||||
|
}
|
||||||
|
museum_df["City_Clean"] = museum_df['City'].replace(city_corrections)
|
||||||
|
museum_df["Visitors_clean"] = museum_df["Visitors"].str.replace(',', '', regex=False)
|
||||||
|
museum_df["Visitors_clean"] = museum_df["Visitors_clean"].str.extract(r'^(\d+)')
|
||||||
|
museum_df["Visitors_clean"] = museum_df["Visitors_clean"].astype(int)
|
||||||
|
museum_df["Visitors"] = museum_df["Visitors_clean"]
|
||||||
|
museum_df = museum_df.drop(columns=["Visitors_clean"])
|
||||||
|
museum_df = _update_city_population(museum_df)
|
||||||
|
return museum_df
|
||||||
|
|
||||||
|
|
||||||
@@ -0,0 +1,158 @@
|
|||||||
|
import os
|
||||||
|
import requests
|
||||||
|
import wikipediaapi
|
||||||
|
from bs4 import BeautifulSoup
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
from museum_analytics import constants
|
||||||
|
from io import StringIO
|
||||||
|
import pandas as pd
|
||||||
|
|
||||||
|
load_dotenv()
|
||||||
|
|
||||||
|
WIKIPEDIA_USER_AGENT = os.getenv("WIKIPEDIA_USER_AGENT")
|
||||||
|
SOUP_PER_SECTION: dict[str, BeautifulSoup] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def get_first_data_table_from_wikipedia(page_title: str) -> pd.DataFrame:
|
||||||
|
"""
|
||||||
|
|
||||||
|
:param page_title:
|
||||||
|
:return:
|
||||||
|
"""
|
||||||
|
|
||||||
|
page_url = f"https://en.wikipedia.org/api/rest_v1/page/html/{page_title}"
|
||||||
|
print(f"Fetching data from Wikipedia: {page_url}")
|
||||||
|
with requests.Session() as session:
|
||||||
|
session.headers["User-Agent"] = WIKIPEDIA_USER_AGENT
|
||||||
|
html = session.get(page_url, timeout=30).text
|
||||||
|
|
||||||
|
tables = pd.read_html(StringIO(html), attrs={"class": "wikitable"})
|
||||||
|
# TODO: add a way to specifically fetch a table
|
||||||
|
df = tables[0]
|
||||||
|
|
||||||
|
if df.iloc[-1].isna().all():
|
||||||
|
df = df.iloc[:-1]
|
||||||
|
return df
|
||||||
|
|
||||||
|
|
||||||
|
def country_table(soup: BeautifulSoup, country: str):
|
||||||
|
for h in soup.find_all(["h2", "h3"]):
|
||||||
|
if h.get_text(strip=True) == country: # exact match, not substring
|
||||||
|
tbl = h.find_next("table", class_="wikitable")
|
||||||
|
return pd.read_html(StringIO(str(tbl)))[0]
|
||||||
|
raise ValueError(f"No section for {country!r}")
|
||||||
|
|
||||||
|
|
||||||
|
def get_country_cities_populations(country: str) -> pd.DataFrame | None:
|
||||||
|
found_section = None
|
||||||
|
for section in constants.CITY_POPULATION_WIKI_SUB_PAGE_NAMES:
|
||||||
|
if country[0] in section:
|
||||||
|
found_section = section
|
||||||
|
break
|
||||||
|
if not found_section:
|
||||||
|
return
|
||||||
|
|
||||||
|
soup = SOUP_PER_SECTION.get(found_section)
|
||||||
|
|
||||||
|
if soup is None:
|
||||||
|
page_title = f"{constants.CITY_POPULATION_WIKI_PAGE_NAME}: {found_section}"
|
||||||
|
|
||||||
|
with requests.Session() as session:
|
||||||
|
session.headers["User-Agent"] = WIKIPEDIA_USER_AGENT
|
||||||
|
response = session.get(
|
||||||
|
"https://en.wikipedia.org/w/api.php", params={
|
||||||
|
"action": "parse", "page": page_title, "prop": "text",
|
||||||
|
"format": "json", "formatversion": 2,
|
||||||
|
}, timeout=30)
|
||||||
|
response.raise_for_status()
|
||||||
|
html = response.json()["parse"]["text"]
|
||||||
|
|
||||||
|
soup = BeautifulSoup(html, "lxml")
|
||||||
|
SOUP_PER_SECTION[found_section] = soup
|
||||||
|
|
||||||
|
dataframe = country_table(soup, country)
|
||||||
|
|
||||||
|
# Population column name includes the year, e.g. "Population (2021)"
|
||||||
|
pop_col = next(c for c in dataframe.columns if str(c).startswith("Population"))
|
||||||
|
dataframe = dataframe.rename(columns={pop_col: "Population"})
|
||||||
|
dataframe["Population"] = pd.to_numeric(
|
||||||
|
dataframe["Population"].astype(str).str.replace(r"\[.*?\]|,", "", regex=True),
|
||||||
|
errors="coerce",
|
||||||
|
)
|
||||||
|
dataframe = dataframe.sort_values("Population", ascending=False)
|
||||||
|
return dataframe
|
||||||
|
|
||||||
|
def get_city_population_data_frame_from_wikipedia(museum_cities_df: pd.DataFrame) -> pd.DataFrame:
|
||||||
|
print("Fetching all population data...")
|
||||||
|
country_list = museum_cities_df["Country"].unique().tolist()
|
||||||
|
data_frames = []
|
||||||
|
for country in country_list:
|
||||||
|
dataframe = get_country_cities_populations(country)
|
||||||
|
if dataframe is None:
|
||||||
|
print(f"Cannot find population data for {country}")
|
||||||
|
continue
|
||||||
|
data_frames.append(dataframe)
|
||||||
|
|
||||||
|
all_cities = pd.concat(data_frames)
|
||||||
|
cities_to_filter = museum_cities_df["City"].unique().tolist()
|
||||||
|
filtered_cities_population = all_cities[all_cities["City"].isin(cities_to_filter)]
|
||||||
|
return filtered_cities_population
|
||||||
|
|
||||||
|
|
||||||
|
def get_city_population_data_frame_from_raw_data(museum_cities_df: pd.DataFrame) -> pd.DataFrame:
|
||||||
|
|
||||||
|
raw_df = pd.read_csv(constants.RAW_POPULATION_DATA_FILE)
|
||||||
|
raw_df = raw_df.drop(["city_ascii", "lat", "lng", "iso2", "iso3", "capital", "id"], axis=1)
|
||||||
|
raw_df["population"] = raw_df["population"].astype("Int64")
|
||||||
|
wrong_washington_condition = (
|
||||||
|
(raw_df["city"] == "Washington")
|
||||||
|
& (raw_df["country"] == "United States")
|
||||||
|
& (raw_df["admin_name"] != "District of Columbia")
|
||||||
|
)
|
||||||
|
cities_population = raw_df[~wrong_washington_condition]
|
||||||
|
cities_population.rename(columns={"population": "Population"},inplace=True)
|
||||||
|
return cities_population
|
||||||
|
|
||||||
|
def update_city_population(museum_df) -> pd.DataFrame:
|
||||||
|
city_pop = get_city_population_data_frame_from_raw_data(museum_df)
|
||||||
|
museum_df = pd.merge(
|
||||||
|
museum_df,
|
||||||
|
city_pop,
|
||||||
|
left_on=["City_Clean", "Country"],
|
||||||
|
right_on=["city", "country"],
|
||||||
|
how="left"
|
||||||
|
)
|
||||||
|
museum_df.drop(columns=["city", "country"], inplace=True)
|
||||||
|
|
||||||
|
return museum_df
|
||||||
|
|
||||||
|
def refresh_museum_data():
|
||||||
|
print("Refreshing museum data...")
|
||||||
|
|
||||||
|
url = constants.DEFAULT_MUSEUM_DATA_SOURCE_URL
|
||||||
|
page = url.rsplit("/", 1)[-1]
|
||||||
|
museum_df = get_first_data_table_from_wikipedia(page_title=page)
|
||||||
|
# print(museum_df.to_string())
|
||||||
|
|
||||||
|
city_corrections = {
|
||||||
|
'Washington, D.C.': 'Washington',
|
||||||
|
'New York City': 'New York',
|
||||||
|
'Vatican City, Rome': 'Vatican City',
|
||||||
|
'London, South Kensington': 'London',
|
||||||
|
}
|
||||||
|
museum_df["City_Clean"] = museum_df['City'].replace(city_corrections)
|
||||||
|
museum_df["Visitors_clean"] = museum_df["Visitors"].str.replace(',', '', regex=False)
|
||||||
|
museum_df["Visitors_clean"] = museum_df["Visitors_clean"].str.extract(r'^(\d+)')
|
||||||
|
museum_df["Visitors_clean"] = museum_df["Visitors_clean"].astype(int)
|
||||||
|
museum_df["Visitors"] = museum_df["Visitors_clean"]
|
||||||
|
museum_df = museum_df.drop(columns=["Visitors_clean"])
|
||||||
|
|
||||||
|
museum_df = update_city_population(museum_df)
|
||||||
|
|
||||||
|
print(museum_df.to_string())
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
refresh_museum_data()
|
||||||
@@ -0,0 +1,290 @@
|
|||||||
|
"""
|
||||||
|
Museum visitors vs. city population: regression model + FastAPI service.
|
||||||
|
|
||||||
|
Usage, with a DataFrame you already have in memory:
|
||||||
|
|
||||||
|
from museum_api import create_app
|
||||||
|
import uvicorn
|
||||||
|
uvicorn.run(create_app(df), port=8000)
|
||||||
|
|
||||||
|
Or from a file:
|
||||||
|
|
||||||
|
MUSEUM_DATA=museums.csv uvicorn museum_api:app --reload
|
||||||
|
|
||||||
|
Expected columns: Name, Visitors, City, Country, Population
|
||||||
|
"""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import os
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Literal
|
||||||
|
import numpy as np
|
||||||
|
import pandas as pd
|
||||||
|
import data_setup
|
||||||
|
import constants
|
||||||
|
from fastapi import FastAPI, HTTPException, Query
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
from sklearn.linear_model import LinearRegression
|
||||||
|
from sklearn.model_selection import KFold, cross_val_score
|
||||||
|
|
||||||
|
|
||||||
|
from dotenv import load_dotenv
|
||||||
|
|
||||||
|
load_dotenv()
|
||||||
|
|
||||||
|
|
||||||
|
REQUIRED_COLUMNS = ["Name", "Visitors", "City", "Country", "Population"]
|
||||||
|
|
||||||
|
|
||||||
|
# --------------------------------------------------------------------------- #
|
||||||
|
# Data + model
|
||||||
|
# --------------------------------------------------------------------------- #
|
||||||
|
def clean(df: pd.DataFrame) -> pd.DataFrame:
|
||||||
|
"""Validate columns, coerce numerics, drop rows that can't be modelled."""
|
||||||
|
missing = set(REQUIRED_COLUMNS) - set(df.columns)
|
||||||
|
if missing:
|
||||||
|
raise ValueError(f"Missing columns: {sorted(missing)}")
|
||||||
|
|
||||||
|
out = df[REQUIRED_COLUMNS].copy()
|
||||||
|
for col in ("Visitors", "Population"):
|
||||||
|
# Handles strings like "1,234,567"
|
||||||
|
out[col] = pd.to_numeric(
|
||||||
|
out[col].astype(str).str.replace(r"[,\s]", "", regex=True),
|
||||||
|
errors="coerce",
|
||||||
|
)
|
||||||
|
out = out.dropna(subset=["Visitors", "Population"])
|
||||||
|
out = out[(out["Visitors"] > 0) & (out["Population"] > 0)]
|
||||||
|
return out.reset_index(drop=True)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class FittedModel:
|
||||||
|
reg: LinearRegression
|
||||||
|
scale: Literal["log", "linear"]
|
||||||
|
data: pd.DataFrame # the rows actually used for fitting
|
||||||
|
n: int
|
||||||
|
r2: float
|
||||||
|
cv_r2: float | None
|
||||||
|
pearson_r: float
|
||||||
|
spearman_rho: float
|
||||||
|
aggregate_by_city: bool
|
||||||
|
predictions: np.ndarray = field(repr=False)
|
||||||
|
|
||||||
|
def _x(self, population: np.ndarray) -> np.ndarray:
|
||||||
|
x = np.asarray(population, dtype=float).reshape(-1, 1)
|
||||||
|
return np.log10(x) if self.scale == "log" else x
|
||||||
|
|
||||||
|
def predict(self, population) -> np.ndarray:
|
||||||
|
y = self.reg.predict(self._x(population))
|
||||||
|
return 10 ** y if self.scale == "log" else y
|
||||||
|
|
||||||
|
@property
|
||||||
|
def slope(self) -> float:
|
||||||
|
return float(self.reg.coef_[0])
|
||||||
|
|
||||||
|
@property
|
||||||
|
def intercept(self) -> float:
|
||||||
|
return float(self.reg.intercept_)
|
||||||
|
|
||||||
|
def summary(self) -> dict:
|
||||||
|
s = {
|
||||||
|
"n_samples": self.n,
|
||||||
|
"scale": self.scale,
|
||||||
|
"aggregate_by_city": self.aggregate_by_city,
|
||||||
|
"slope": self.slope,
|
||||||
|
"intercept": self.intercept,
|
||||||
|
"r2": self.r2,
|
||||||
|
"cv_r2_5fold": self.cv_r2,
|
||||||
|
"pearson_r": self.pearson_r,
|
||||||
|
"spearman_rho": self.spearman_rho,
|
||||||
|
}
|
||||||
|
if self.scale == "log":
|
||||||
|
s["equation"] = (
|
||||||
|
f"visitors = {10 ** self.intercept:,.1f} * population^{self.slope:.3f}"
|
||||||
|
)
|
||||||
|
s["interpretation"] = (
|
||||||
|
f"A 10x larger city is associated with "
|
||||||
|
f"{10 ** self.slope:.2f}x the visitors."
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
s["equation"] = f"visitors = {self.slope:.4f} * population + {self.intercept:,.0f}"
|
||||||
|
s["interpretation"] = (
|
||||||
|
f"Each additional 1M residents is associated with "
|
||||||
|
f"{self.slope * 1e6:,.0f} more visitors."
|
||||||
|
)
|
||||||
|
return s
|
||||||
|
|
||||||
|
|
||||||
|
def fit_model(
|
||||||
|
df: pd.DataFrame,
|
||||||
|
scale: Literal["log", "linear"] = "log",
|
||||||
|
aggregate_by_city: bool = False,
|
||||||
|
) -> FittedModel:
|
||||||
|
"""
|
||||||
|
Fit visitors ~ population.
|
||||||
|
|
||||||
|
scale="log" (default) fits log10(visitors) ~ log10(population). Both
|
||||||
|
variables are heavily right-skewed, so a log-log fit is usually far
|
||||||
|
better behaved than a linear one and the slope reads as an elasticity.
|
||||||
|
|
||||||
|
aggregate_by_city=True sums visitors per (City, Country) first, so a
|
||||||
|
city with many museums counts once rather than once per museum.
|
||||||
|
"""
|
||||||
|
data = clean(df)
|
||||||
|
if aggregate_by_city:
|
||||||
|
data = (
|
||||||
|
data.groupby(["City", "Country"], as_index=False)
|
||||||
|
.agg(Visitors=("Visitors", "sum"),
|
||||||
|
Population=("Population", "first"),
|
||||||
|
Name=("Name", lambda s: ", ".join(s)))
|
||||||
|
)
|
||||||
|
if len(data) < 3:
|
||||||
|
raise ValueError(f"Need at least 3 usable rows, got {len(data)}")
|
||||||
|
|
||||||
|
x_raw = data["Population"].to_numpy(float)
|
||||||
|
y_raw = data["Visitors"].to_numpy(float)
|
||||||
|
X = (np.log10(x_raw) if scale == "log" else x_raw).reshape(-1, 1)
|
||||||
|
y = np.log10(y_raw) if scale == "log" else y_raw
|
||||||
|
|
||||||
|
reg = LinearRegression().fit(X, y)
|
||||||
|
|
||||||
|
cv_r2 = None
|
||||||
|
if len(data) >= 10:
|
||||||
|
cv = KFold(n_splits=5, shuffle=True, random_state=0)
|
||||||
|
cv_r2 = float(cross_val_score(LinearRegression(), X, y, cv=cv, scoring="r2").mean())
|
||||||
|
|
||||||
|
xs, ys = pd.Series(X.ravel()), pd.Series(y)
|
||||||
|
model = FittedModel(
|
||||||
|
reg=reg,
|
||||||
|
scale=scale,
|
||||||
|
data=data,
|
||||||
|
n=len(data),
|
||||||
|
r2=float(reg.score(X, y)),
|
||||||
|
cv_r2=cv_r2,
|
||||||
|
pearson_r=float(xs.corr(ys)),
|
||||||
|
# Spearman = Pearson on ranks (avoids a scipy dependency)
|
||||||
|
spearman_rho=float(pd.Series(x_raw).rank().corr(pd.Series(y_raw).rank())),
|
||||||
|
aggregate_by_city=aggregate_by_city,
|
||||||
|
predictions=np.empty(0),
|
||||||
|
)
|
||||||
|
model.predictions = model.predict(x_raw)
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
# --------------------------------------------------------------------------- #
|
||||||
|
# API schemas
|
||||||
|
# --------------------------------------------------------------------------- #
|
||||||
|
class MuseumRecord(BaseModel):
|
||||||
|
Name: str
|
||||||
|
Visitors: float = Field(gt=0)
|
||||||
|
City: str
|
||||||
|
Country: str
|
||||||
|
Population: float = Field(gt=0)
|
||||||
|
|
||||||
|
|
||||||
|
class TrainRequest(BaseModel):
|
||||||
|
records: list[MuseumRecord] = Field(min_length=3)
|
||||||
|
scale: Literal["log", "linear"] = "log"
|
||||||
|
aggregate_by_city: bool = False
|
||||||
|
|
||||||
|
|
||||||
|
class PredictRequest(BaseModel):
|
||||||
|
populations: list[float] = Field(min_length=1)
|
||||||
|
|
||||||
|
|
||||||
|
class Prediction(BaseModel):
|
||||||
|
population: float
|
||||||
|
predicted_visitors: float
|
||||||
|
|
||||||
|
|
||||||
|
# --------------------------------------------------------------------------- #
|
||||||
|
# App factory
|
||||||
|
# --------------------------------------------------------------------------- #
|
||||||
|
def create_app(
|
||||||
|
df: pd.DataFrame | None = None,
|
||||||
|
scale: Literal["log", "linear"] = "log",
|
||||||
|
aggregate_by_city: bool = False,
|
||||||
|
) -> FastAPI:
|
||||||
|
app = FastAPI(title="Museum Visitors Regression", version="1.0")
|
||||||
|
app.state.model = fit_model(df, scale, aggregate_by_city) if df is not None else None
|
||||||
|
|
||||||
|
def get_model() -> FittedModel:
|
||||||
|
if app.state.model is None:
|
||||||
|
raise HTTPException(503, "No model trained yet. POST /train first.")
|
||||||
|
return app.state.model
|
||||||
|
|
||||||
|
@app.get("/health")
|
||||||
|
def health():
|
||||||
|
return {"status": "ok", "model_loaded": app.state.model is not None}
|
||||||
|
|
||||||
|
@app.get("/model")
|
||||||
|
def model_summary():
|
||||||
|
return get_model().summary()
|
||||||
|
|
||||||
|
@app.post("/train")
|
||||||
|
def train(req: TrainRequest):
|
||||||
|
df_new = pd.DataFrame([r.model_dump() for r in req.records])
|
||||||
|
try:
|
||||||
|
app.state.model = fit_model(df_new, req.scale, req.aggregate_by_city)
|
||||||
|
except ValueError as e:
|
||||||
|
raise HTTPException(422, str(e))
|
||||||
|
return app.state.model.summary()
|
||||||
|
|
||||||
|
@app.get("/predict", response_model=Prediction)
|
||||||
|
def predict_one(population: float = Query(gt=0)):
|
||||||
|
m = get_model()
|
||||||
|
return Prediction(population=population,
|
||||||
|
predicted_visitors=float(m.predict([population])[0]))
|
||||||
|
|
||||||
|
@app.post("/predict", response_model=list[Prediction])
|
||||||
|
def predict_many(req: PredictRequest):
|
||||||
|
if any(p <= 0 for p in req.populations):
|
||||||
|
raise HTTPException(422, "Populations must be > 0")
|
||||||
|
m = get_model()
|
||||||
|
preds = m.predict(req.populations)
|
||||||
|
return [Prediction(population=p, predicted_visitors=float(v))
|
||||||
|
for p, v in zip(req.populations, preds)]
|
||||||
|
|
||||||
|
@app.get("/residuals")
|
||||||
|
def residuals(
|
||||||
|
top: int = Query(10, ge=1, le=500),
|
||||||
|
order: Literal["over", "under"] = "over",
|
||||||
|
):
|
||||||
|
"""Museums that most over/under-perform what their city size predicts."""
|
||||||
|
m = get_model()
|
||||||
|
d = m.data.copy()
|
||||||
|
d["Predicted"] = m.predictions
|
||||||
|
d["Ratio"] = d["Visitors"] / d["Predicted"] # >1 = outperforms city size
|
||||||
|
d = d.sort_values("Ratio", ascending=(order == "under")).head(top)
|
||||||
|
return d[["Name", "City", "Country", "Population",
|
||||||
|
"Visitors", "Predicted", "Ratio"]].to_dict(orient="records")
|
||||||
|
|
||||||
|
return app
|
||||||
|
|
||||||
|
|
||||||
|
def _load_data() -> pd.DataFrame | None:
|
||||||
|
relative_path = constants.MUSEUM_DATA_FILE
|
||||||
|
|
||||||
|
absolute_path = os.path.abspath(relative_path)
|
||||||
|
|
||||||
|
if not os.path.exists(absolute_path):
|
||||||
|
print("Fetching and conforming museum data and city population...")
|
||||||
|
museum_data = data_setup.get_museum_data()
|
||||||
|
museum_data.to_csv(absolute_path)
|
||||||
|
else:
|
||||||
|
museum_data = pd.read_csv(absolute_path)
|
||||||
|
|
||||||
|
return museum_data
|
||||||
|
|
||||||
|
|
||||||
|
# Module-level app so `uvicorn museum_api:app` works.
|
||||||
|
app = create_app(
|
||||||
|
_load_data(),
|
||||||
|
scale=os.environ.get("MUSEUM_SCALE", "log"), # type: ignore[arg-type]
|
||||||
|
aggregate_by_city=os.environ.get("MUSEUM_AGG_CITY", "0") == "1",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
import uvicorn
|
||||||
|
uvicorn.run(app, host="0.0.0.0", port=8000)
|
||||||
Reference in new issue
Block a user