Spaces:
Runtime error
Runtime error
| import json | |
| import re | |
| import os | |
| from io import BytesIO | |
| from pathlib import Path | |
| from functools import lru_cache | |
| from typing import Dict, List, Optional, Tuple | |
| import numpy as np | |
| import pandas as pd | |
| import plotly.express as px | |
| import plotly.graph_objects as go | |
| import gradio as gr | |
| from datasets import load_dataset | |
| # --- Optional token (HuggingFace) --- | |
| HF_TOKEN = os.environ.get("HF_TOKEN") | |
| # --- Required libs (fail early with clear message) --- | |
| try: | |
| import geopandas as gpd | |
| except Exception as e: | |
| raise ImportError("Please install GeoPandas: pip install geopandas shapely fiona pyproj") from e | |
| try: | |
| from wordcloud import WordCloud, STOPWORDS | |
| except Exception as e: | |
| raise ImportError("Please install wordcloud: pip install wordcloud") from e | |
| # Pillow is used for the word cloud image type | |
| from PIL import Image | |
| # ================== CONFIG ================== | |
| HERE = Path(__file__).parent | |
| FILES: Dict[str, str] = { | |
| "Hebei": "hebeicombined.json", | |
| "Shanxi": "shanxicombined.json", | |
| "Shaanxi": "shaanxicombined.json", | |
| "Shandong": "shandongcombined.json", | |
| "Henan": "henancombined.json", | |
| "Hubei": "hubeicombined.json", | |
| "Jiangsu": "jiangsucombined.json", | |
| "Anhui": "anhuicombined.json", | |
| "Zhejiang": "zhejiangcombined.json", | |
| "Hunan": "hunancombined.json", | |
| "Jiangxi": "jiangxicombined.json", | |
| "Fujian": "fujiancombined.json", | |
| "Guangdong": "guangdongcombined.json", | |
| } | |
| GEO_PATH = "cn.geojson" | |
| VMIN, VMAX = 0.0, 10.0 | |
| # ================== STYLING ================== | |
| BG_WINDOW = "#0b1020" | |
| CARD_BG = "rgba(255,255,255,0.06)" | |
| BORDER = "rgba(255,255,255,0.18)" | |
| TEXT = "#000000" | |
| ACCENT = "#8DC3F9" | |
| FONT_STACK = "-apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, 'Helvetica Neue', Arial, sans-serif" | |
| def apply_layout(fig: go.Figure, title: str) -> go.Figure: | |
| fig.update_layout( | |
| template="presentation", | |
| title=dict(text=title, x=0.01, xanchor="left", font=dict(size=20, color=TEXT)), | |
| paper_bgcolor="rgba(0,0,0,0)", | |
| plot_bgcolor="rgba(0,0,0,0)", | |
| font=dict(family=FONT_STACK, size=14, color=TEXT), | |
| margin=dict(l=80, r=30, t=64, b=40), | |
| ) | |
| return fig | |
| # ================== NAME NORMALIZATION (DRY) ================== | |
| ALIASES = { | |
| "Inner Mongolia Autonomous Region": "Inner Mongolia", | |
| "Xinjiang Uygur Autonomous Region": "Xinjiang", | |
| "Tibet Autonomous Region": "Tibet", | |
| "Guangxi Zhuang Autonomous Region": "Guangxi", | |
| "Ningxia Hui Autonomous Region": "Ningxia", | |
| "Nei Mongol": "Inner Mongolia", | |
| "Xizang": "Tibet", | |
| "Shaan Xi": "Shaanxi", | |
| "Hong Kong SAR": "Hong Kong", | |
| "Macao SAR": "Macau", | |
| "Macau": "Macau", | |
| } | |
| _SUFFIX_RX = re.compile(r"\s+(Province|Autonomous Region|Region|City|Prefecture)$", flags=re.I) | |
| def norm_name(s: str) -> str: | |
| s = str(s or "").strip() | |
| s = _SUFFIX_RX.sub("", s) | |
| s = (s.replace("Xinjiang Uygur", "Xinjiang") | |
| .replace("Tibet Autonomous", "Tibet") | |
| .replace("Guangxi Zhuang", "Guangxi") | |
| .replace("Ningxia Hui", "Ningxia") | |
| .replace("Shaan Xi", "Shaanxi")) | |
| return ALIASES.get(s, s) | |
| # ================== DATA LOADING ================== | |
| def load_from_individual_json(files: Dict[str, str]) -> pd.DataFrame: | |
| rows = [] | |
| for prov, fp in files.items(): | |
| p = HERE / fp | |
| if not p.exists(): | |
| continue | |
| with p.open("r", encoding="utf-8") as f: | |
| data = json.load(f) | |
| for r in data: | |
| y, s = r.get("year"), r.get("score") | |
| if y is None or s is None: | |
| continue | |
| rows.append({ | |
| "province": norm_name(r.get("province", prov)), | |
| "year": int(y), | |
| "score": float(s), | |
| "rationale": r.get("rationale"), | |
| "confidence": r.get("confidence"), | |
| "key_evidence": r.get("key_evidence"), | |
| }) | |
| df = pd.DataFrame(rows) | |
| if df.empty: | |
| raise ValueError("No data loaded from JSONs — check FILES paths.") | |
| if "key_evidence" in df.columns: | |
| df["key_evidence"] = df["key_evidence"].apply(lambda x: "; ".join(x) if isinstance(x, list) else x) | |
| return df | |
| def load_df() -> pd.DataFrame: | |
| try: | |
| ds = load_dataset("csv", data_files="hf://datasets/zliang/LLMdrought/droughtdata.csv", token=HF_TOKEN) | |
| df = ds["train"].to_pandas() | |
| df["province"] = df["province"].map(norm_name) | |
| return df | |
| except Exception: | |
| df = load_from_individual_json(FILES) | |
| return df | |
| DF = load_df() | |
| PROVINCES = sorted(DF["province"].unique().tolist()) | |
| YEAR_MIN, YEAR_MAX = int(DF["year"].min()), int(DF["year"].max()) | |
| YEARS = np.arange(YEAR_MIN, YEAR_MAX + 1) | |
| PROV_ORDER = [p for p in FILES.keys() if norm_name(p) in PROVINCES] or PROVINCES | |
| # ---------- Fast matrices & cumulative stats ---------- | |
| def score_matrix() -> Tuple[np.ndarray, List[str], np.ndarray]: | |
| """Dense province×year matrix aligned to PROV_ORDER × YEARS.""" | |
| M = ( | |
| DF.pivot_table(index="province", columns="year", values="score", aggfunc="mean") | |
| .reindex(index=PROV_ORDER).reindex(columns=YEARS) | |
| ) | |
| return M.to_numpy(dtype=np.float64), M.index.tolist(), M.columns.to_numpy(dtype=int) | |
| def raw_matrix_df() -> pd.DataFrame: | |
| M, idx, cols = score_matrix() | |
| return pd.DataFrame(M, index=idx, columns=cols) | |
| def cumstats(): | |
| """Cumulative sum & count (NaN-aware) for O(1) range means.""" | |
| M, idx, cols = score_matrix() | |
| valid = ~np.isnan(M) | |
| Mz = np.nan_to_num(M, nan=0.0) | |
| S = np.cumsum(Mz, axis=1) # cumulative sum | |
| C = np.cumsum(valid.astype(np.int32), 1) # cumulative count | |
| return S, C, idx, cols | |
| def range_mean_by_province(y0: int, y1: int, provinces: List[str]) -> pd.DataFrame: | |
| y0, y1 = int(y0), int(y1) | |
| if y0 > y1: y0, y1 = y1, y0 | |
| y0 = max(YEAR_MIN, min(y0, YEAR_MAX)) | |
| y1 = max(YEAR_MIN, min(y1, YEAR_MAX)) | |
| S, C, idx, cols = cumstats() | |
| colpos = {int(y): i for i, y in enumerate(cols)} | |
| i0, i1 = colpos[y0], colpos[y1] | |
| sum_win = S[:, i1] - (S[:, i0-1] if i0 > 0 else 0) | |
| cnt_win = C[:, i1] - (C[:, i0-1] if i0 > 0 else 0) | |
| mean_win = np.divide(sum_win, cnt_win, out=np.full_like(sum_win, np.nan), where=cnt_win > 0) | |
| out = pd.DataFrame({"province": idx, "drought_mean": mean_win}) | |
| if provinces: | |
| sel = set(norm_name(p) for p in provinces) | |
| out = out[out["province"].isin(sel)] | |
| out["name"] = out["province"] | |
| return out | |
| def rolling_matrix(window: int) -> pd.DataFrame: | |
| """Fast NaN-aware rolling mean along years using cumulative sums. | |
| Keeps the same number of columns as the base matrix (min_periods=1).""" | |
| if window <= 1: | |
| return raw_matrix_df() | |
| base = raw_matrix_df() | |
| M = base.to_numpy(copy=True) # (rows, N) | |
| rows, N = M.shape | |
| valid = ~np.isnan(M) | |
| Mz = np.nan_to_num(M, nan=0.0) | |
| w = int(window) | |
| pad = w - 1 | |
| # Left-pad so early columns use fewer than w terms (min_periods=1 behavior) | |
| Mz_pad = np.pad(Mz, ((0, 0), (pad, 0))) | |
| V_pad = np.pad(valid.astype(np.float64), ((0, 0), (pad, 0))) | |
| # Cumulative sums over padded arrays | |
| CS_M = np.cumsum(Mz_pad, axis=1) # shape (rows, N+pad) | |
| CS_V = np.cumsum(V_pad, axis=1) | |
| # For each output column j (0..N-1), window ends at index end = pad + j | |
| # and starts at start = end - (w-1). Use CS[end] - CS[start-1] (with 0 when start==0). | |
| end_idx = pad + np.arange(N) # shape (N,) | |
| start_idx = end_idx - (w - 1) # shape (N,) | |
| r = np.arange(rows)[:, None] # row indices for broadcasting | |
| # sums over windows | |
| sum_win = CS_M[r, end_idx[None, :]] - np.where( | |
| start_idx[None, :] > 0, | |
| CS_M[r, (start_idx - 1)[None, :]], | |
| 0.0 | |
| ) | |
| cnt_win = CS_V[r, end_idx[None, :]] - np.where( | |
| start_idx[None, :] > 0, | |
| CS_V[r, (start_idx - 1)[None, :]], | |
| 0.0 | |
| ) | |
| R = np.divide(sum_win, cnt_win, out=np.full_like(sum_win, np.nan), where=cnt_win > 0) | |
| return pd.DataFrame(R, index=base.index, columns=base.columns) | |
| # ================== GEO ================== | |
| def load_geodata(path: str, simplify_tolerance: float = 0.05) -> "gpd.GeoDataFrame": | |
| gdf = gpd.read_file(path) | |
| if gdf.empty: | |
| raise ValueError(f"Geo file at {path} is empty or unreadable.") | |
| try: | |
| if gdf.crs and gdf.crs.to_string() not in ("EPSG:4326", "WGS84"): | |
| gdf = gdf.to_crs(epsg=4326) | |
| except Exception: | |
| pass | |
| candidates = [c for c in ["name","NAME","NAME_1","NAME_EN","Province","province","NL_NAME_1","Name"] if c in gdf.columns] | |
| source = candidates[0] if candidates else next((c for c in gdf.columns if gdf[c].dtype == object), None) | |
| gdf = gdf.copy() | |
| gdf["name"] = gdf[source].map(norm_name) if source else "Unknown" | |
| # simplify geometry once to cut payload size | |
| try: | |
| gdf["geometry"] = gdf.geometry.simplify(simplify_tolerance, preserve_topology=True) | |
| except Exception: | |
| pass | |
| return gdf | |
| GDF = load_geodata(GEO_PATH) | |
| # Build a lean GeoJSON dict once for Plotly | |
| GEOJSON = json.loads(GDF.to_json()) | |
| GEO_FEATURE_KEY = "properties.name" | |
| # ================== UTILS ================== | |
| def clamp_years(y0: int, y1: int) -> Tuple[int, int]: | |
| y0, y1 = int(y0), int(y1) | |
| if y0 > y1: y0, y1 = y1, y0 | |
| y0 = max(YEAR_MIN, min(y0, YEAR_MAX)) | |
| y1 = max(YEAR_MIN, min(y1, YEAR_MAX)) | |
| return y0, y1 | |
| def is_english(text: str, min_ascii_ratio: float = 0.8) -> bool: | |
| if not text: return False | |
| ascii_chars = sum(1 for ch in text if ord(ch) < 128) | |
| return (ascii_chars / max(1, len(text))) >= min_ascii_ratio | |
| # ================== FIGURES ================== | |
| def apply_choropleth_colorscale_name(name: str) -> str: | |
| m = { | |
| "ylorrd": "YlOrRd", | |
| "viridis": "Viridis", | |
| "plasma": "Plasma", | |
| "inferno": "Inferno", | |
| "cividis": "Cividis", | |
| "blues": "Blues", | |
| } | |
| return m.get((name or "").lower(), "YlOrRd") | |
| def build_choropleth(y0: int, y1: int, provinces: List[str], cmap: str) -> Tuple[go.Figure, str]: | |
| y0, y1 = clamp_years(y0, y1) | |
| stat = range_mean_by_province(y0, y1, provinces or []) | |
| matched = int(stat["drought_mean"].notna().sum()) | |
| total = len(stat) | |
| status = f"Matched {matched}/{total} provinces." | |
| fig = px.choropleth( | |
| stat, | |
| geojson=GEOJSON, # prebuilt once | |
| locations="name", | |
| featureidkey=GEO_FEATURE_KEY, | |
| color="drought_mean", | |
| color_continuous_scale=apply_choropleth_colorscale_name(cmap), | |
| range_color=(VMIN, VMAX), | |
| ) | |
| fig.update_traces(marker_line_color="black", marker_line_width=0.6) | |
| fig.update_geos(fitbounds="locations", visible=True, showcountries=True, showocean=False, showland=False, showlakes=False) | |
| fig.update_coloraxes(colorbar=dict(title=f"Mean score ({y0}–{y1})", thickness=14, len=0.7)) | |
| return apply_layout(fig, f"China drought mean — {y0}–{y1}"), status | |
| def composite_series(provinces: List[str], y0: int, y1: int, roll: int, z_mult: float = 1.96) -> go.Figure: | |
| sel = [norm_name(p) for p in (provinces or PROVINCES)] | |
| y0, y1 = clamp_years(y0, y1) | |
| cols = [y for y in YEARS if y0 <= y <= y1] | |
| mat = rolling_matrix(roll) if roll and roll > 1 else raw_matrix_df() | |
| sub = mat.loc[[p for p in sel if p in mat.index], cols] | |
| mean = sub.mean(axis=0, skipna=True).to_numpy() | |
| std = sub.std(axis=0, ddof=1, skipna=True).to_numpy() | |
| n = sub.count(axis=0).to_numpy() | |
| sem = np.divide(std, np.sqrt(np.maximum(n, 1)), out=np.zeros_like(std), where=n > 0) | |
| upper = mean + z_mult * sem | |
| lower = mean - z_mult * sem | |
| fig = go.Figure() | |
| fig.add_trace(go.Scatter(x=cols, y=upper, line=dict(width=0), hoverinfo="skip", showlegend=False)) | |
| fig.add_trace(go.Scatter(x=cols, y=lower, fill="tonexty", fillcolor="rgba(10,132,255,0.22)", line=dict(width=0), hoverinfo="skip", showlegend=False)) | |
| fig.add_trace(go.Scatter(x=cols, y=mean, mode="lines", name=f"Composite mean (n={len(sub.index)})", line=dict(width=3, color=ACCENT), hovertemplate="Year %{x}<br>Mean %{y:.2f}<extra></extra>")) | |
| fig.update_xaxes(title="Year"); fig.update_yaxes(title="Drought score") | |
| return apply_layout(fig, f"Composite time series — {roll or 1}-year smoothing, 95% band") | |
| def _plotly_colorscale(name: str) -> Tuple[str, bool]: | |
| name = (name or "").strip().lower() | |
| if name == "ylorrd": return "YlOrRd", False | |
| if name == "viridis": return "Viridis", False | |
| if name == "plasma": return "Plasma", False | |
| if name == "inferno": return "Inferno", False | |
| if name == "cividis": return "Cividis", False | |
| if name == "blues": return "Blues", False | |
| if name == "rdbu_r": return "RdBu", True # reversed matplotlib-style alias | |
| return "YlOrRd", False | |
| def heatmap_figure(y0: int, y1: int, use_roll5: bool = True, cmap: str = "YlOrRd") -> go.Figure: | |
| y0, y1 = clamp_years(y0, y1) | |
| mat = rolling_matrix(5) if use_roll5 else raw_matrix_df() | |
| cols = [y for y in YEARS if y0 <= y <= y1] | |
| M = mat[cols] | |
| Z = M.to_numpy() | |
| Z = np.where(np.isnan(Z), None, Z) # Plotly-friendly | |
| colorscale, reversescale = _plotly_colorscale(cmap) | |
| fig = go.Figure( | |
| data=go.Heatmap( | |
| z=Z, | |
| x=cols, | |
| y=M.index.tolist(), | |
| colorscale=colorscale, | |
| reversescale=reversescale, | |
| zmin=VMIN if colorscale == "YlOrRd" else None, | |
| zmax=VMAX if colorscale == "YlOrRd" else None, | |
| colorbar=dict(title="Drought score" + (" (roll-5)" if use_roll5 else "")), | |
| hovertemplate="Year %{x}<br>Province %{y}<br>Score %{z:.2f}<extra></extra>", | |
| ) | |
| ) | |
| fig.update_xaxes(title="Year") | |
| fig.update_yaxes(title="Province", autorange="reversed") | |
| return apply_layout(fig, f"Province × Year heatmap — {'roll-5' if use_roll5 else 'raw'}") | |
| # --------- Word Cloud ---------- | |
| def english_text_blob(prov_key: str, y0: int, y1: int, min_ascii_ratio: float) -> str: | |
| y0, y1 = clamp_years(y0, y1) | |
| d = DF[(DF["year"] >= y0) & (DF["year"] <= y1)] | |
| if prov_key: | |
| sel = set(prov_key.split("|")) | |
| d = d[d["province"].isin(sel)] | |
| texts = [] | |
| for rat in d["rationale"].dropna().astype(str): | |
| if is_english(rat, min_ascii_ratio=min_ascii_ratio): | |
| texts.append(rat) | |
| return " ".join(texts) if texts else "no english rationale text found in selection and range" | |
| def make_wordcloud_image( | |
| provinces: List[str], y0: int, y1: int, | |
| bg_color: str = "white", min_ascii_ratio: float = 0.8, | |
| extra_stopwords: Optional[str] = None, width: int = 1280, height: int = 800, | |
| ) -> Image.Image: | |
| y0, y1 = clamp_years(y0, y1) | |
| prov_key = "|".join(sorted(norm_name(p) for p in (provinces or []))) | |
| all_text = english_text_blob(prov_key, int(y0), int(y1), float(min_ascii_ratio)) | |
| sw = set(STOPWORDS) | { | |
| "indicating","reported","county","includes","year","include","due","source", | |
| "score","confidence","counter","evidence","and","no","counties","in","with", | |
| "the","of","to","for","a","is","that","it","this","as","on","are","was","by", | |
| "an","at","be","from","which","or","not,","have","has","but","its","their","may", | |
| "also","more","other","such","these","than","most","some","one","two","three", | |
| "four","five","six","seven","eight","nine","ten","indicating","suggest, indicate" | |
| } | |
| if extra_stopwords: | |
| sw |= {w.strip().lower() for w in re.split(r"[,;\s]+", extra_stopwords) if w.strip()} | |
| wc = WordCloud(width=width, height=height, background_color=bg_color, collocations=False, prefer_horizontal=0.9, stopwords=sw).generate(all_text) | |
| return wc.to_image() | |
| # ================== UI ================== | |
| def build_interface() -> gr.Blocks: | |
| theme = gr.themes.Soft() | |
| palettes = ["YlOrRd", "Viridis", "Plasma", "Inferno", "Cividis", "Blues"] | |
| with gr.Blocks(title="China Drought index derived from historical documents by LLM", theme=theme) as demo: | |
| gr.Markdown("## China Drought index derived from historical documents by LLM") | |
| start_default = max(1638, YEAR_MIN) | |
| end_default = min(1644, YEAR_MAX) | |
| with gr.Tabs(): | |
| # ============ Introduction TAB ============ | |
| with gr.Tab("Introduction"): | |
| gr.Markdown( | |
| "# Introduction\n" | |
| "Imagine flipping through old Chinese chronicles and discovering hidden clues about the past climate! " | |
| "These historical records don’t just tell stories of emperors, they also preserve vivid descriptions of past climate events. " | |
| "By tapping into these records, we can trace long-term climate trends and patterns across history. \n\n" | |
| "In this study, we introduce an ordinal-scale drought index derived from these records by a large language model (LLM). " | |
| "Our analysis spans thirteen Chinese provinces from 1368 to 1911, covering the Ming and Qing dynasties. " | |
| "Use the tabs to explore different visualizations and find your favorite way to see the drought trend:\n\n" | |
| "🗺️ **Map**: Explore geospatial patterns of drought across chosen periods \n\n" | |
| "📈 **Composite Time Series**: Discover national average drought trends over time \n\n" | |
| "🔥 **Province Heatmap**: Dive into province-specific drought trends over time \n\n" | |
| "☁️ **Rationale Word Cloud**: See the keywords highlighted by the LLM when evaluating drought severity from descriptions" | |
| ) | |
| # ============ MAP TAB ============ | |
| with gr.Tab("🗺️ Map"): | |
| with gr.Row(): | |
| provinces_inp = gr.CheckboxGroup(label="Provinces", choices=PROVINCES, value=PROVINCES) | |
| with gr.Column(): | |
| start_year = gr.Slider(label="Start year", minimum=YEAR_MIN, maximum=YEAR_MAX, value=start_default, step=1) | |
| end_year = gr.Slider(label="End year", minimum=YEAR_MIN, maximum=YEAR_MAX, value=end_default, step=1) | |
| with gr.Accordion("Style", open=False): | |
| palette_dd = gr.Dropdown(label="Palette", choices=palettes, value="YlOrRd") | |
| status_md = gr.Markdown() | |
| map_plot = gr.Plot(label="Choropleth (mean over selected years)", show_label=False) | |
| # ============ TIME SERIES TAB ============ | |
| with gr.Tab("📈 Composite Time Series"): | |
| with gr.Row(): | |
| provinces_inp_ts = gr.CheckboxGroup(label="Provinces", choices=PROVINCES, value=PROVINCES) | |
| with gr.Column(): | |
| range_start = gr.Slider(label="Start year", minimum=YEAR_MIN, maximum=YEAR_MAX, value=YEAR_MIN, step=1) | |
| range_end = gr.Slider(label="End year", minimum=YEAR_MIN, maximum=YEAR_MAX, value=YEAR_MAX, step=1) | |
| series_roll = gr.Slider(label="Rolling mean (years)", minimum=0, maximum=30, value=5, step=1) | |
| series_plot = gr.Plot(label="Composite mean with 95% confidence band", show_label=False) | |
| # ============ HEATMAP TAB ============ | |
| with gr.Tab("🔥 Province Heatmap"): | |
| with gr.Row(): | |
| with gr.Column(): | |
| hm_start = gr.Slider(label="Start year", minimum=YEAR_MIN, maximum=YEAR_MAX, value=YEAR_MIN, step=1) | |
| hm_end = gr.Slider(label="End year", minimum=YEAR_MIN, maximum=YEAR_MAX, value=YEAR_MAX, step=1) | |
| with gr.Column(): | |
| hm_roll = gr.Checkbox(label="Use 5-year rolling mean", value=True) | |
| hm_cmap = gr.Dropdown(label="Colormap", choices=["YlOrRd","RdBu_r","viridis","plasma","inferno","Cividis","Blues"], value="YlOrRd") | |
| heatmap_plot = gr.Plot(label="Province × year heatmap", show_label=False) | |
| # ============ WORD CLOUD TAB ============ | |
| with gr.Tab("☁️ Rationale Word Cloud"): | |
| with gr.Row(): | |
| provinces_inp_wc = gr.CheckboxGroup(label="Provinces", choices=PROVINCES, value=PROVINCES) | |
| with gr.Column(): | |
| wc_start = gr.Slider(label="Start year", minimum=YEAR_MIN, maximum=YEAR_MAX, value=YEAR_MIN, step=1) | |
| wc_end = gr.Slider(label="End year", minimum=YEAR_MIN, maximum=YEAR_MAX, value=YEAR_MAX, step=1) | |
| with gr.Accordion("Style", open=False): | |
| wc_bg = gr.Dropdown(label="Background", choices=["white","black"], value="white") | |
| wc_ascii = gr.Slider(label="English filter (min ASCII ratio)", minimum=0.5, maximum=1.0, value=0.8, step=0.05) | |
| wc_stop = gr.Textbox(label="Extra stopwords (comma/space separated)", placeholder="e.g. drought famine river province …") | |
| wc_btn = gr.Button("Generate word cloud", variant="primary") | |
| wc_img = gr.Image(label="Word Cloud", type="pil") | |
| # ====== Callbacks ====== | |
| def on_update_map(prov, y0, y1, cmap): | |
| fig, status = build_choropleth(int(y0), int(y1), prov or [], cmap) | |
| return fig, status | |
| def on_update_series(prov, y0, y1, roll): | |
| return composite_series(prov or [], int(y0), int(y1), int(roll)) | |
| def on_update_heatmap(y0, y1, roll5, cmap): | |
| return heatmap_figure(int(y0), int(y1), use_roll5=bool(roll5), cmap=cmap) | |
| def on_update_wc(prov, y0, y1, bg, ascii_ratio, extra_sw): | |
| return make_wordcloud_image(prov or [], int(y0), int(y1), bg_color=bg, min_ascii_ratio=float(ascii_ratio), extra_stopwords=extra_sw) | |
| # ---- Bindings with release() when available to avoid redraw spam ---- | |
| def bind_release_or_change(comp, fn, inputs, outputs): | |
| """Use .release() if present (Gradio v4), else .change().""" | |
| if hasattr(comp, "release"): | |
| comp.release(fn=fn, inputs=inputs, outputs=outputs) | |
| else: | |
| comp.change(fn=fn, inputs=inputs, outputs=outputs) | |
| # Map | |
| provinces_inp.change(fn=on_update_map, inputs=[provinces_inp, start_year, end_year, palette_dd], outputs=[map_plot, status_md]) | |
| bind_release_or_change(start_year, on_update_map, [provinces_inp, start_year, end_year, palette_dd], [map_plot, status_md]) | |
| bind_release_or_change(end_year, on_update_map, [provinces_inp, start_year, end_year, palette_dd], [map_plot, status_md]) | |
| palette_dd.change(fn=on_update_map, inputs=[provinces_inp, start_year, end_year, palette_dd], outputs=[map_plot, status_md]) | |
| demo.load(on_update_map, [provinces_inp, start_year, end_year, palette_dd], [map_plot, status_md]) | |
| # Series | |
| provinces_inp_ts.change(fn=on_update_series, inputs=[provinces_inp_ts, range_start, range_end, series_roll], outputs=[series_plot]) | |
| bind_release_or_change(range_start, on_update_series, [provinces_inp_ts, range_start, range_end, series_roll], [series_plot]) | |
| bind_release_or_change(range_end, on_update_series, [provinces_inp_ts, range_start, range_end, series_roll], [series_plot]) | |
| bind_release_or_change(series_roll, on_update_series, [provinces_inp_ts, range_start, range_end, series_roll], [series_plot]) | |
| demo.load(on_update_series, [provinces_inp_ts, range_start, range_end, series_roll], [series_plot]) | |
| # Heatmap | |
| bind_release_or_change(hm_start, on_update_heatmap, [hm_start, hm_end, hm_roll, hm_cmap], [heatmap_plot]) | |
| bind_release_or_change(hm_end, on_update_heatmap, [hm_start, hm_end, hm_roll, hm_cmap], [heatmap_plot]) | |
| hm_roll.change(fn=on_update_heatmap, inputs=[hm_start, hm_end, hm_roll, hm_cmap], outputs=[heatmap_plot]) | |
| hm_cmap.change(fn=on_update_heatmap, inputs=[hm_start, hm_end, hm_roll, hm_cmap], outputs=[heatmap_plot]) | |
| demo.load(on_update_heatmap, [hm_start, hm_end, hm_roll, hm_cmap], [heatmap_plot]) | |
| # Word cloud — button only (no live changes to avoid heavy CPU on each tweak) | |
| wc_btn.click(fn=on_update_wc, inputs=[provinces_inp_wc, wc_start, wc_end, wc_bg, wc_ascii, wc_stop], outputs=[wc_img]) | |
| return demo | |
| # ================== MAIN ================== | |
| def main(): | |
| app = build_interface() | |
| # app.launch(server_name="0.0.0.0", server_port=7860) # optional for containers | |
| app.launch() | |
| if __name__ == "__main__": | |
| main() | |