LiveHouse-TS / scripts /export_livehouse_tables.py
ziyuzhou02's picture
Deploy GitHub main 3feb6cda1511
e317359 verified
Raw History Blame Contribute Delete
7.32 kB
#!/usr/bin/env python3
"""Export LiveHouse per-dataset paper tables from live aggregate results."""
from __future__ import annotations
import argparse
import csv
import math
from pathlib import Path
MODELS = [
"Chronos-2",
"TiRex",
"TimesFM-2.5",
"Toto-1.0",
"Moirai-2.0",
"Chronos-Bolt",
"TabPFN-TS",
"Sundial",
"Moving-Average",
"ETS",
"ARIMA",
"Seasonal-Naive",
]
MODEL_SEPARATOR_AFTER = "Sundial"
DATASETS = [
("BTC", "binance_btcusdt_1h_close/1h/short"),
("PM2.5", "open_meteo_shanghai_air_quality_hourly_pm2_5/1h/short"),
("Quake", "usgs_earthquake_all_week_hourly_earthquake_count/1h/short"),
("Potomac", "usgs_potomac_iv_usgs_00060/15min/short"),
("Water", "noaa_coops_sf_water_level_water_level/6min/short"),
("Wave", "noaa_ndbc_46013_realtime_wave_height/10min/short"),
("T2M", "nasa_power_shanghai_hourly_T2M/1h/short"),
("KSFO", "nws_ksfo_observations_temperature/1h/short"),
("Temp2m", "open_meteo_shanghai_hourly_temperature_2m/1h/short"),
("Wiki", "wikimedia_time_series_hourly_pageviews/1D/short"),
]
METRICS = [
(
"RMSE",
r"Per-dataset RMSE ($\downarrow$) on the online benchmark.",
"tab_perdataset_rmse",
),
(
"MAPE",
r"Per-dataset MAPE ($\downarrow$) on the online benchmark.",
"tab_perdataset_mape",
),
(
"CRPS",
r"Per-dataset CRPS ($\downarrow$) on the online benchmark.",
"tab_perdataset_crps",
),
(
"Stability",
r"Per-dataset Temporal Stability ($\downarrow$) on the online benchmark, computed from release-level MSE histories.",
"tab_perdataset_stability",
),
(
"Improvement",
r"Per-dataset Improvement ($\downarrow$; more negative is better) on the online benchmark, computed as Kendall $\tau$ on release-level MSE histories.",
"tab_perdataset_improvement",
),
]
HIGHLIGHT_NOTE = (
r"In each column the best available model is in \textcolor{red}{\textbf{red bold}} "
r"and the second best is in \textcolor{blue}{\underline{blue underline}}."
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--aggregate",
type=Path,
default=Path("space/results/aggregates/live_results_by_dataset.csv"),
help="Path to live_results_by_dataset.csv.",
)
parser.add_argument(
"--output",
type=Path,
default=Path("paper/livehouse_up_to_now_tables.tex"),
help="LaTeX output path.",
)
parser.add_argument(
"--fill-single-release-zero",
action="store_true",
help="For Stability/Improvement only, write 0.000 when n_releases < 2.",
)
return parser.parse_args()
def parse_float(value: str | None) -> float | None:
if value is None or value == "":
return None
try:
parsed = float(value)
except ValueError:
return None
if not math.isfinite(parsed):
return None
return parsed
def load_rows(path: Path) -> dict[tuple[str, str], dict[str, str]]:
rows: dict[tuple[str, str], dict[str, str]] = {}
with path.open(newline="", encoding="utf-8") as handle:
for row in csv.DictReader(handle):
rows[(row["model"], row["dataset"])] = row
return rows
def release_count(row: dict[str, str] | None) -> int:
if not row:
return 0
try:
return int(float(row.get("n_releases", "0") or 0))
except ValueError:
return 0
def metric_value(
rows: dict[tuple[str, str], dict[str, str]],
model: str,
dataset: str,
metric: str,
fill_single_release_zero: bool,
) -> float | None:
row = rows.get((model, dataset))
if row is None:
return None
if metric in {"Stability", "Improvement"} and release_count(row) < 2:
return 0.0 if fill_single_release_zero else None
return parse_float(row.get(metric))
def format_value(value: float | None) -> str:
if value is None:
return "--"
abs_value = abs(value)
if abs_value >= 1000 or (0 < abs_value < 0.001):
return f"{value:.3e}"
return f"{value:.3f}"
def highlight(formatted: str, rank: int | None) -> str:
if rank == 1:
return rf"\textcolor{{red}}{{\textbf{{{formatted}}}}}"
if rank == 2:
return rf"\textcolor{{blue}}{{\underline{{{formatted}}}}}"
return formatted
def ranks_for_column(values: dict[str, float | None]) -> dict[str, int | None]:
ranked = [
(model, value)
for model, value in values.items()
if value is not None and math.isfinite(value)
]
ranked.sort(key=lambda item: (item[1], MODELS.index(item[0])))
ranks = {model: None for model in values}
if ranked:
ranks[ranked[0][0]] = 1
if len(ranked) > 1:
ranks[ranked[1][0]] = 2
return ranks
def caption_for(metric: str, base_caption: str, fill_single_release_zero: bool) -> str:
caption = base_caption
if metric in {"Stability", "Improvement"}:
if fill_single_release_zero:
caption += " Cells with fewer than two releases are provisionally assigned 0."
else:
caption += " Cells with fewer than two releases are shown as --."
return f"{caption} {HIGHLIGHT_NOTE}"
def render_metric_table(
rows: dict[tuple[str, str], dict[str, str]],
metric: str,
caption: str,
label: str,
fill_single_release_zero: bool,
) -> str:
values_by_dataset = {
dataset_id: {
model: metric_value(rows, model, dataset_id, metric, fill_single_release_zero)
for model in MODELS
}
for _, dataset_id in DATASETS
}
ranks_by_dataset = {
dataset_id: ranks_for_column(values)
for dataset_id, values in values_by_dataset.items()
}
lines = [
r"\begin{table*}[!t]",
r"\centering",
rf"\caption{{{caption_for(metric, caption, fill_single_release_zero)}}}",
rf"\label{{{label}}}",
r"\resizebox{0.8 \linewidth}{!}{%",
r"\begin{tabular}{l rrrrrrrrrr}",
r"\toprule",
"Model & " + " & ".join(name for name, _ in DATASETS) + r" \\",
r"\midrule",
]
for model in MODELS:
cells = []
for _, dataset_id in DATASETS:
value = values_by_dataset[dataset_id][model]
formatted = format_value(value)
rank = ranks_by_dataset[dataset_id][model]
cells.append(highlight(formatted, rank))
lines.append(model + " & " + " & ".join(cells) + r" \\")
if model == MODEL_SEPARATOR_AFTER:
lines.append(r"\midrule")
lines.extend(
[
r"\bottomrule",
r"\end{tabular}}",
r"\end{table*}",
]
)
return "\n".join(lines)
def main() -> None:
args = parse_args()
rows = load_rows(args.aggregate)
tables = [
render_metric_table(rows, metric, caption, label, args.fill_single_release_zero)
for metric, caption, label in METRICS
]
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text("\n\n".join(tables) + "\n", encoding="utf-8")
print(f"Wrote {args.output}")
if __name__ == "__main__":
main()