Spaces:
Sleeping
Sleeping
Download app.py from stardust-coder/MinDeMoDeMo: direct link, hf CLI and curl.
- Browser
- Download file 29 kB
-
https://huggingface.co/spaces/stardust-coder/MinDeMoDeMo/resolve/main/app.py
- Command line
-
hf download hf://spaces/stardust-coder/MinDeMoDeMo/app.py
-
curl -L -o app.py https://huggingface.co/spaces/stardust-coder/MinDeMoDeMo/resolve/main/app.py
29 kB
| # -*- coding: utf-8 -*- | |
| import json | |
| import os | |
| import tempfile | |
| import zipfile | |
| from itertools import combinations, product | |
| from pathlib import Path | |
| import gradio as gr | |
| import networkx as nx | |
| import numpy as np | |
| import pandas as pd | |
| import plotly.express as px | |
| import plotly.graph_objects as go | |
| import mindemo_v4 as mindemo | |
| # ============================================================ | |
| # Presets | |
| # ============================================================ | |
| DATA_PRESETS = { | |
| "なし": None, | |
| "synthetic_iid_small": {"N": 200, "d": 5, "kind": "iid", "seed": 0}, | |
| "synthetic_timeseries_small": {"N": 300, "d": 6, "kind": "timeseries", "seed": 1}, | |
| # Same basic data setup as the penguin example in the mindemo repository: | |
| # Adelie penguins, four continuous body measurements, and sex. | |
| "palmerpenguins_adelie": {"kind": "palmerpenguins"}, | |
| } | |
| # h must contain terms involving at least two coordinates; pure one-variable terms | |
| # cancel under the coordinate-swap construction used by Besag PLE. | |
| H_PRESETS = { | |
| "pairwise_product": {"degree": 2, "max_interaction_order": 2}, | |
| "pairwise_polynomial_degree_3": {"degree": 3, "max_interaction_order": 2}, | |
| "pairwise_polynomial_degree_4": {"degree": 4, "max_interaction_order": 2}, | |
| "degree_3_with_threeway": {"degree": 3, "max_interaction_order": 3}, | |
| } | |
| ESTIMATORS = { | |
| "Besag PLE (SGD系)": "sgd", | |
| "Besag PLE (full batch)": "fullbatch", | |
| "CLE (full batch / MCMC)": "cle", | |
| } | |
| OPTIMIZERS = ["sgd", "adam", "adamw", "rmsprop", "adagrad"] | |
| # ============================================================ | |
| # Data utilities | |
| # ============================================================ | |
| def make_synthetic_data(preset_name: str) -> pd.DataFrame: | |
| preset = DATA_PRESETS[preset_name] | |
| rng = np.random.default_rng(preset["seed"]) | |
| N, d = preset["N"], preset["d"] | |
| if preset["kind"] == "iid": | |
| # Correlated Gaussian sample so pairwise dependence is visible. | |
| z = rng.normal(size=(N, d)) | |
| X = z.copy() | |
| for j in range(1, d): | |
| X[:, j] = 0.35 * X[:, j - 1] + np.sqrt(1 - 0.35**2) * z[:, j] | |
| else: | |
| X = np.zeros((N, d)) | |
| A = rng.normal(scale=0.15, size=(d, d)) | |
| np.fill_diagonal(A, 0.3) | |
| noise = rng.normal(scale=0.5, size=(N, d)) | |
| for t in range(1, N): | |
| X[t] = X[t - 1] @ A.T + noise[t] | |
| return pd.DataFrame(X, columns=[f"x{j}" for j in range(d)]) | |
| def make_palmerpenguins_adelie_data() -> pd.DataFrame: | |
| """Load the Adelie subset used by the mindemo penguin example. | |
| The original example studies dependence among four continuous measurements | |
| and sex. The estimator/UI downstream expects numeric columns, so sex is | |
| encoded as female=0 and male=1. Rows missing any of the five variables are | |
| removed. | |
| """ | |
| try: | |
| from palmerpenguins import load_penguins | |
| except ImportError as e: | |
| raise gr.Error( | |
| "palmerpenguins がインストールされていません。" | |
| " `pip install palmerpenguins` を実行してください。" | |
| ) from e | |
| penguins = load_penguins() | |
| columns = [ | |
| "bill_length_mm", | |
| "bill_depth_mm", | |
| "flipper_length_mm", | |
| "body_mass_g", | |
| "sex", | |
| ] | |
| df = penguins.loc[penguins["species"] == "Adelie", columns].copy() | |
| df = df.dropna(subset=columns) | |
| #正規化を入れておく ################## | |
| continuous_cols = [ | |
| "bill_length_mm", | |
| "bill_depth_mm", | |
| "flipper_length_mm", | |
| "body_mass_g", | |
| ] | |
| df[continuous_cols] = ( | |
| df[continuous_cols] - df[continuous_cols].mean() | |
| ) / df[continuous_cols].std() | |
| ################################# | |
| sex_map = {"female": 0.0, "male": 1.0} | |
| df["sex"] = df["sex"].astype(str).str.lower().map(sex_map) | |
| df = df.dropna(subset=["sex"]) | |
| return validate_data(df.reset_index(drop=True)) | |
| def load_csv_or_preset(csv_file, preset_name: str) -> pd.DataFrame: | |
| if csv_file is not None: | |
| csv_path = csv_file if isinstance(csv_file, str) else csv_file.name | |
| return validate_data(pd.read_csv(csv_path)) | |
| if preset_name != "なし": | |
| preset = DATA_PRESETS[preset_name] | |
| if preset.get("kind") == "palmerpenguins": | |
| return make_palmerpenguins_adelie_data() | |
| return make_synthetic_data(preset_name) | |
| raise gr.Error("CSVをアップロードするか、データプリセットを選んでください。") | |
| def validate_data(df: pd.DataFrame) -> pd.DataFrame: | |
| if df.empty: | |
| raise gr.Error("データが空です。") | |
| numeric_df = df.select_dtypes(include=[np.number]) | |
| if numeric_df.shape[1] == 0: | |
| raise gr.Error("数値列がありません。N x d の数値CSVを指定してください。") | |
| if numeric_df.isna().any().any(): | |
| raise gr.Error("欠損値があります。事前に補完または削除してください。") | |
| if numeric_df.shape[0] < 2: | |
| raise gr.Error("PLEには少なくとも2観測が必要です。") | |
| return numeric_df | |
| def data_preview_plot(df: pd.DataFrame, data_mode: str, max_points: int = 1000): | |
| if df.empty: | |
| return go.Figure() | |
| if data_mode == "i.i.d.": | |
| plot_df = df.iloc[:max_points, :6].copy() | |
| fig = px.scatter_matrix( | |
| plot_df, | |
| dimensions=list(plot_df.columns), | |
| title="Data preview pair plot", | |
| ) | |
| fig.update_traces(diagonal_visible=True, showupperhalf=False, marker=dict(size=4, opacity=0.55)) | |
| fig.update_layout(height=720, dragmode="select") | |
| return fig | |
| plot_df = df.iloc[:max_points, :12].copy() | |
| plot_df.insert(0, "index", np.arange(len(plot_df))) | |
| long_df = plot_df.melt(id_vars="index", var_name="variable", value_name="value") | |
| fig = px.line(long_df, x="index", y="value", color="variable", title="Data preview time-series plot") | |
| fig.update_layout(height=420, legend_title_text="variables") | |
| return fig | |
| def parse_custom_h_json(custom_h_json: str) -> dict: | |
| if not custom_h_json.strip(): | |
| return {} | |
| try: | |
| obj = json.loads(custom_h_json) | |
| except json.JSONDecodeError as e: | |
| raise gr.Error(f"関数 h のカスタムJSONが不正です: {e}") | |
| if not isinstance(obj, dict): | |
| raise gr.Error("関数 h のカスタム指定はJSON objectにしてください。") | |
| return obj | |
| # ============================================================ | |
| # Canonical statistic h | |
| # ============================================================ | |
| def _positive_compositions(total: int, parts: int): | |
| """Yield tuples of positive integers of length parts summing to total.""" | |
| if parts == 1: | |
| yield (total,) | |
| return | |
| for first in range(1, total - parts + 2): | |
| for rest in _positive_compositions(total - first, parts - 1): | |
| yield (first,) + rest | |
| def build_h_function(var_names: list[str], h_config: dict): | |
| """ | |
| Build polynomial interaction statistics. | |
| Every term contains >=2 distinct variables. This avoids statistics that | |
| vanish identically in the coordinate-swap contrast used by Besag PLE. | |
| """ | |
| d = len(var_names) | |
| degree = int(h_config.get("degree", 2)) | |
| max_order = int(h_config.get("max_interaction_order", 2)) | |
| if degree < 2: | |
| raise gr.Error("h の degree は2以上にしてください。") | |
| if max_order < 2: | |
| raise gr.Error("max_interaction_order は2以上にしてください。") | |
| max_order = min(max_order, d, degree) | |
| terms = [] | |
| labels = [] | |
| supports = [] | |
| for order in range(2, max_order + 1): | |
| for idxs in combinations(range(d), order): | |
| for total_degree in range(order, degree + 1): | |
| for exps in _positive_compositions(total_degree, order): | |
| terms.append((idxs, exps)) | |
| pieces = [] | |
| for idx, exp in zip(idxs, exps): | |
| pieces.append(var_names[idx] if exp == 1 else f"{var_names[idx]}^{exp}") | |
| labels.append(" * ".join(pieces)) | |
| supports.append(tuple(idxs)) | |
| if not terms: | |
| raise gr.Error("h の項が生成されませんでした。設定を確認してください。") | |
| def h(x): | |
| x = np.asarray(x, dtype=float) | |
| out = np.empty(len(terms), dtype=float) | |
| for k, (idxs, exps) in enumerate(terms): | |
| value = 1.0 | |
| for idx, exp in zip(idxs, exps): | |
| value *= x[idx] ** exp | |
| out[k] = value | |
| return out | |
| return h, labels, supports | |
| # ============================================================ | |
| # mindemo_v4 adapter | |
| # ============================================================ | |
| def run_mindemo_method( | |
| X: np.ndarray, | |
| estimator_label: str, | |
| optimizer: str, | |
| lr: float, | |
| max_iter: int, | |
| tol: float, | |
| batch_size: int, | |
| stepsize_decay: float, | |
| l2_penalty: float, | |
| use_rpj: bool, | |
| avg_start: int, | |
| fullbatch_l1: bool, | |
| fullbatch_C: float, | |
| h, | |
| random_seed: int, | |
| cle_L: int, | |
| cle_burnin: int, | |
| cle_thin: int, | |
| ): | |
| X_list = X.tolist() | |
| mode = ESTIMATORS[estimator_label] | |
| if mode == "sgd": | |
| theta, rpj_count = mindemo.mindemo_Besag_SGD( | |
| X=X_list, | |
| h=h, | |
| max_iter=int(max_iter), | |
| tol=float(tol), | |
| learningrate=float(lr), | |
| stepsize_decay=float(stepsize_decay), | |
| batch_size=int(batch_size), | |
| seed=int(random_seed), | |
| use_rpj=bool(use_rpj), | |
| avg_start=int(avg_start), | |
| return_avg=True, | |
| optimizer=str(optimizer), | |
| l2_penalty=float(l2_penalty), | |
| verbose_every=0, | |
| ) | |
| return np.asarray(theta, dtype=float), { | |
| "estimator": estimator_label, | |
| "optimizer": optimizer, | |
| "learningrate": float(lr), | |
| "max_iter": int(max_iter), | |
| "tol": float(tol), | |
| "batch_size": int(batch_size), | |
| "stepsize_decay": float(stepsize_decay), | |
| "l2_penalty": float(l2_penalty), | |
| "use_rpj": bool(use_rpj), | |
| "avg_start": int(avg_start), | |
| "rpj_count": None if np.isinf(rpj_count) else int(rpj_count), | |
| } | |
| if mode == "fullbatch": | |
| if fullbatch_C <= 0: | |
| raise gr.Error("full-batch の C は0より大きくしてください。") | |
| theta = mindemo.mindemo_Besag_fullbatch( | |
| X=X_list, | |
| h=h, | |
| max_iter=int(max_iter), | |
| tol=float(tol), | |
| sparseop=bool(fullbatch_l1), | |
| sparsepen=float(fullbatch_C), | |
| ) | |
| return np.asarray(theta, dtype=float), { | |
| "estimator": estimator_label, | |
| "optimizer": "scikit-learn LogisticRegression (mindemo_v4)", | |
| "max_iter": int(max_iter), | |
| "tol": float(tol), | |
| "l1_enabled": bool(fullbatch_l1), | |
| "C": float(fullbatch_C), | |
| } | |
| # CLE: use the implementation and defaults from mindemo_v4, with exposed MCMC controls. | |
| theta_list, res_list, L_list, _ = mindemo.mindemo_CLE_fullbatch( | |
| X=X_list, | |
| h=h, | |
| L=int(cle_L), | |
| burnin=int(cle_burnin), | |
| thin=int(cle_thin), | |
| max_iter=int(max_iter), | |
| tol=float(tol), | |
| detailop=False, | |
| init=None, | |
| ) | |
| if len(theta_list) == 0: | |
| raise gr.Error("CLE が推定値を返しませんでした。") | |
| theta = np.asarray(theta_list[-1], dtype=float) | |
| return theta, { | |
| "estimator": estimator_label, | |
| "max_iter": int(max_iter), | |
| "tol": float(tol), | |
| "L_initial": int(cle_L), | |
| "burnin": int(cle_burnin), | |
| "thin": int(cle_thin), | |
| "iterations_returned": len(theta_list), | |
| "residual_history": [float(x) for x in res_list], | |
| "L_history": [int(x) for x in L_list], | |
| } | |
| # ============================================================ | |
| # Result shaping / visualization | |
| # ============================================================ | |
| def coefficient_table(theta, labels, supports, var_names): | |
| rows = [] | |
| for k, (coef, label, support) in enumerate(zip(theta, labels, supports)): | |
| rows.append({ | |
| "term_index": k, | |
| "term": label, | |
| "variables": ", ".join(var_names[i] for i in support), | |
| "order": len(support), | |
| "estimate": float(coef), | |
| "abs_estimate": abs(float(coef)), | |
| }) | |
| return pd.DataFrame(rows).sort_values("abs_estimate", ascending=False).reset_index(drop=True) | |
| def pairwise_aggregate(theta, supports, var_names): | |
| # Multiple polynomial terms may correspond to the same pair. Aggregate by L2 | |
| # magnitude while retaining the sign of the largest-magnitude term. | |
| buckets = {} | |
| for coef, support in zip(theta, supports): | |
| if len(support) != 2: | |
| continue | |
| key = tuple(sorted(support)) | |
| buckets.setdefault(key, []).append(float(coef)) | |
| rows = [] | |
| for (i, j), vals in buckets.items(): | |
| vals = np.asarray(vals, dtype=float) | |
| idx = int(np.argmax(np.abs(vals))) | |
| signed_strength = float(np.sign(vals[idx]) * np.sqrt(np.sum(vals**2))) | |
| rows.append({ | |
| "var1": var_names[i], | |
| "var2": var_names[j], | |
| "aggregate_strength": signed_strength, | |
| "abs_strength": abs(signed_strength), | |
| "n_terms": int(len(vals)), | |
| }) | |
| if not rows: | |
| return pd.DataFrame(columns=["var1", "var2", "aggregate_strength", "abs_strength", "n_terms"]) | |
| return pd.DataFrame(rows).sort_values("abs_strength", ascending=False).reset_index(drop=True) | |
| def coefficient_plot(coef_df: pd.DataFrame, top_n: int = 40): | |
| if coef_df.empty: | |
| return go.Figure() | |
| plot_df = coef_df.head(top_n).iloc[::-1] | |
| fig = px.bar(plot_df, x="estimate", y="term", orientation="h", title=f"Estimated coefficients (top {min(top_n, len(coef_df))})") | |
| fig.update_layout(height=max(420, 22 * len(plot_df))) | |
| return fig | |
| def graph_plot( | |
| pair_df: pd.DataFrame, | |
| var_names: list[str], | |
| threshold: float, | |
| layout_preset: str | None = None, | |
| ): | |
| G = nx.Graph() | |
| G.add_nodes_from(var_names) | |
| for _, row in pair_df.iterrows(): | |
| if float(row["abs_strength"]) >= threshold: | |
| G.add_edge(row["var1"], row["var2"], weight=float(row["aggregate_strength"])) | |
| if len(G.edges) == 0: | |
| fig = go.Figure() | |
| fig.update_layout( | |
| title=f"Estimated pairwise structure | no edges above threshold={threshold}", | |
| height=520, | |
| xaxis=dict(visible=False), yaxis=dict(visible=False), | |
| annotations=[dict(text="No pairwise edges above threshold", x=0.5, y=0.5, showarrow=False, font=dict(size=18))], | |
| ) | |
| return fig | |
| # Use a fixed, paper-comparison-friendly layout for the Palmer Penguins | |
| # preset. Other datasets retain the automatic spring layout. | |
| # Match the geometry of Figure S.1 in the paper as closely as practical: | |
| # 1 = bill length -> top | |
| # 2 = bill depth -> bottom | |
| # 5 = sex -> left-middle | |
| # 4 = body mass -> right-middle | |
| # 3 = flipper length -> farther right, horizontally aligned with body mass | |
| penguin_pos = { | |
| "bill_length_mm": (0.0, 1.0), | |
| "bill_depth_mm": (0.0, -1.0), | |
| "sex": (-0.85, 0.0), | |
| "body_mass_g": (0.85, 0.0), | |
| "flipper_length_mm": (1.75, 0.0), | |
| } | |
| if ( | |
| layout_preset == "palmerpenguins_adelie" | |
| and set(var_names) == set(penguin_pos) | |
| ): | |
| pos = {name: penguin_pos[name] for name in G.nodes} | |
| else: | |
| pos = nx.spring_layout(G, seed=42, k=0.8) | |
| edge_x, edge_y, annotations = [], [], [] | |
| for src, dst, data in G.edges(data=True): | |
| x0, y0 = pos[src] | |
| x1, y1 = pos[dst] | |
| edge_x += [x0, x1, None] | |
| edge_y += [y0, y1, None] | |
| # Show the estimated pairwise coefficient at the midpoint of each edge. | |
| # Explicitly setting ``text`` also avoids Plotly's default "new text" label. | |
| weight = float(data["weight"]) | |
| annotations.append(dict( | |
| x=(x0 + x1) / 2.0, | |
| y=(y0 + y1) / 2.0, | |
| xref="x", | |
| yref="y", | |
| text=f"{weight:.3f}", | |
| showarrow=False, | |
| font=dict(size=13), | |
| bgcolor="rgba(255,255,255,0.80)", | |
| borderpad=2, | |
| )) | |
| edge_trace = go.Scatter(x=edge_x, y=edge_y, line=dict(width=1.5), hoverinfo="none", mode="lines") | |
| node_trace = go.Scatter( | |
| x=[pos[n][0] for n in G.nodes], y=[pos[n][1] for n in G.nodes], | |
| mode="markers+text", text=list(G.nodes), textposition="top center", | |
| marker=dict(size=18, line=dict(width=1)), hoverinfo="text", | |
| ) | |
| fig = go.Figure(data=[edge_trace, node_trace]) | |
| fig.update_layout( | |
| title=f"Estimated pairwise structure | threshold={threshold}", | |
| showlegend=False, height=520, xaxis=dict(visible=False), yaxis=dict(visible=False), | |
| annotations=annotations, margin=dict(l=20, r=20, t=60, b=20), | |
| ) | |
| return fig | |
| def diagnostics_plot(meta: dict): | |
| residuals = meta.get("residual_history") | |
| if residuals: | |
| df = pd.DataFrame({"iteration": np.arange(len(residuals)), "residual": residuals}) | |
| fig = px.line(df, x="iteration", y="residual", title="CLE residual history") | |
| fig.update_layout(height=360) | |
| return fig | |
| fig = go.Figure() | |
| fig.update_layout( | |
| title="Optimization diagnostics", | |
| height=360, | |
| xaxis=dict(visible=False), yaxis=dict(visible=False), | |
| annotations=[dict( | |
| text="mindemo_v4 の PLE 関数は loss history を返さないため、擬似的な loss は表示しません。", | |
| x=0.5, y=0.5, showarrow=False, | |
| )], | |
| ) | |
| return fig | |
| def export_results(coef_df, pair_df, meta): | |
| tmpdir = Path(tempfile.mkdtemp()) | |
| out_path = tmpdir / "estimation_results.zip" | |
| coef_path = tmpdir / "coefficients.csv" | |
| pair_path = tmpdir / "pairwise_summary.csv" | |
| meta_path = tmpdir / "meta.json" | |
| coef_df.to_csv(coef_path, index=False) | |
| pair_df.to_csv(pair_path, index=False) | |
| meta_path.write_text(json.dumps(meta, ensure_ascii=False, indent=2), encoding="utf-8") | |
| with zipfile.ZipFile(out_path, "w") as zf: | |
| zf.write(coef_path, arcname="coefficients.csv") | |
| zf.write(pair_path, arcname="pairwise_summary.csv") | |
| zf.write(meta_path, arcname="meta.json") | |
| return str(out_path) | |
| # ============================================================ | |
| # App callbacks | |
| # ============================================================ | |
| def preview_data(csv_file, data_preset, data_mode): | |
| df = load_csv_or_preset(csv_file, data_preset) | |
| N, d = df.shape | |
| shape_text = f"### 実データサイズ: **{N} × {d}** (`N={N}, d={d}`)" | |
| return df.head(20), data_preview_plot(df, data_mode), shape_text | |
| def run_experiment( | |
| csv_file, data_preset, data_mode, | |
| estimator_label, optimizer, lr, max_iter, tol, batch_size, stepsize_decay, | |
| l2_penalty, use_rpj, avg_start, | |
| fullbatch_l1, fullbatch_C, | |
| h_preset, custom_h_json, | |
| graph_threshold, random_seed, | |
| cle_L, cle_burnin, cle_thin, | |
| ): | |
| df = load_csv_or_preset(csv_file, data_preset) | |
| X = df.to_numpy(dtype=float) | |
| var_names = list(df.columns) | |
| if data_mode != "i.i.d.": | |
| # The supplied mindemo estimators assume an i.i.d. sample. We allow the | |
| # run for experimentation, but state the assumption explicitly. | |
| assumption_warning = "⚠️ mindemo_v4 の推定理論は i.i.d. 標本を前提とします。時系列モードでは行をそのまま標本として渡しています。" | |
| else: | |
| assumption_warning = "" | |
| h_config = dict(H_PRESETS[h_preset]) | |
| h_config.update(parse_custom_h_json(custom_h_json)) | |
| h, labels, supports = build_h_function(var_names, h_config) | |
| theta, algo_meta = run_mindemo_method( | |
| X=X, | |
| estimator_label=estimator_label, | |
| optimizer=optimizer, | |
| lr=float(lr), | |
| max_iter=int(max_iter), | |
| tol=float(tol), | |
| batch_size=int(batch_size), | |
| stepsize_decay=float(stepsize_decay), | |
| l2_penalty=float(l2_penalty), | |
| use_rpj=bool(use_rpj), | |
| avg_start=int(avg_start), | |
| fullbatch_l1=bool(fullbatch_l1), | |
| fullbatch_C=float(fullbatch_C), | |
| h=h, | |
| random_seed=int(random_seed), | |
| cle_L=int(cle_L), | |
| cle_burnin=int(cle_burnin), | |
| cle_thin=int(cle_thin), | |
| ) | |
| if theta.shape != (len(labels),): | |
| raise gr.Error(f"theta の shape が不正です。期待: {(len(labels),)}, 実際: {theta.shape}") | |
| coef_df = coefficient_table(theta, labels, supports, var_names) | |
| pair_df = pairwise_aggregate(theta, supports, var_names) | |
| coef_fig = coefficient_plot(coef_df) | |
| # Only use the fixed penguin layout when the built-in penguin preset is | |
| # actually the data source. If a CSV is uploaded, keep the generic layout. | |
| layout_preset = data_preset if csv_file is None else None | |
| graph_fig = graph_plot( | |
| pair_df, | |
| var_names, | |
| float(graph_threshold), | |
| layout_preset=layout_preset, | |
| ) | |
| diag_fig = diagnostics_plot(algo_meta) | |
| meta = { | |
| **algo_meta, | |
| "N": int(df.shape[0]), | |
| "d": int(df.shape[1]), | |
| "columns": var_names, | |
| "data_mode": data_mode, | |
| "h_preset": h_preset, | |
| "h_config": h_config, | |
| "K": int(len(labels)), | |
| "graph_threshold": float(graph_threshold), | |
| "random_seed": int(random_seed), | |
| } | |
| result_zip = export_results(coef_df, pair_df, meta) | |
| n_edges = int((pair_df["abs_strength"] >= float(graph_threshold)).sum()) if not pair_df.empty else 0 | |
| summary = f""" | |
| ### 実験完了 | |
| - データ形状: **N={df.shape[0]}, d={df.shape[1]}** | |
| - 推定法: **{estimator_label}** | |
| - optimizer: **{optimizer if ESTIMATORS[estimator_label] == 'sgd' else 'N/A'}** | |
| - canonical statistics: **{h_preset}** | |
| - パラメータ次元: **K={len(labels)}** | |
| - 閾値以上の pairwise edge 数: **{n_edges}** | |
| {assumption_warning} | |
| """ | |
| return summary, coef_fig, graph_fig, diag_fig, coef_df, pair_df, result_zip | |
| # ============================================================ | |
| # Gradio UI | |
| # ============================================================ | |
| CSS = """ | |
| #sidebar { border-right: 1px solid var(--border-color-primary); padding-right: 16px; } | |
| .small-note { font-size: 0.9em; color: var(--body-text-color-subdued); } | |
| """ | |
| with gr.Blocks(title="MIDM Estimation Demo") as demo: | |
| gr.Markdown( | |
| """ | |
| # MinDEMO DEMO | |
| Minimum Information Dependence Model (MinDeMo) を用いた多変量データの従属性推定を実行するデモ(DEMO)アプリです。 | |
| `mindemo_v4.py` の Besag PLE / CLE 実装を使用します。CSVは `N x d` の数値データを想定します。 | |
| """ | |
| ) | |
| gr.Markdown("---") | |
| with gr.Row(): | |
| with gr.Column(scale=1, elem_id="sidebar"): | |
| gr.Markdown("## 実験設定") | |
| with gr.Accordion("1. データ", open=False): | |
| csv_file = gr.File(label="CSVアップロード", file_types=[".csv"], type="filepath") | |
| data_preset = gr.Dropdown(choices=list(DATA_PRESETS.keys()), value="synthetic_iid_small", label="データプリセット") | |
| preview_btn = gr.Button("データをプレビュー") | |
| with gr.Accordion("2. データ仮定", open=False): | |
| data_mode = gr.Radio(choices=["i.i.d.", "時系列"], value="i.i.d.", label="データの種類") | |
| gr.Markdown("<div class='small-note'>mindemo_v4 の推定理論は i.i.d. 標本を前提とします。</div>") | |
| with gr.Accordion("3. 推定法", open=True): | |
| estimator = gr.Dropdown(choices=list(ESTIMATORS.keys()), value="Besag PLE (SGD系)", label="推定法") | |
| optimizer = gr.Dropdown(choices=OPTIMIZERS, value="adam", label="optimizer(SGD系のみ)") | |
| lr = gr.Number(value=0.01, label="learning rate(SGD系のみ)") | |
| max_iter = gr.Slider(minimum=10, maximum=20000, value=2000, step=10, label="max_iter") | |
| tol = gr.Number(value=1e-5, label="tolerance") | |
| batch_size = gr.Slider(minimum=1, maximum=512, value=32, step=1, label="batch_size(SGD系のみ)") | |
| stepsize_decay = gr.Number(value=0.0, label="stepsize_decay(SGD系のみ)") | |
| random_seed = gr.Number(value=0, precision=0, label="random seed") | |
| with gr.Accordion("4. SGD安定化 / 正則化", open=False): | |
| l2_penalty = gr.Number(value=0.0, minimum=0.0, label="L2 penalty(SGD系)") | |
| use_rpj = gr.Checkbox(value=True, label="Ruppert–Polyak–Juditsky averaging") | |
| avg_start = gr.Number(value=1000, precision=0, label="RPJ avg_start") | |
| fullbatch_l1 = gr.Checkbox(value=False, label="L1 penalty(full batchのみ)") | |
| fullbatch_C = gr.Number(value=1.0, minimum=1e-8, label="full-batch sparsepen / sklearn C(小さいほど強いL1)") | |
| with gr.Accordion("5. canonical statistic h", open=True): | |
| h_preset = gr.Dropdown(choices=list(H_PRESETS.keys()), value="pairwise_product", label="h のプリセット") | |
| custom_h_json = gr.Code( | |
| value="", language="json", label="カスタム h 設定 JSON(任意)", lines=5, | |
| ) | |
| gr.Markdown( | |
| """<div class='small-note'>例: {"degree": 3, "max_interaction_order": 3}<br> | |
| 各項は2変数以上を含む多項式相互作用として生成します。</div>""" | |
| ) | |
| with gr.Accordion("6. CLE MCMC設定", open=False): | |
| cle_L = gr.Number(value=1000, precision=0, label="L") | |
| cle_burnin = gr.Number(value=100, precision=0, label="burnin") | |
| cle_thin = gr.Number(value=10, precision=0, label="thin") | |
| with gr.Accordion("7. グラフ可視化", open=False): | |
| graph_threshold = gr.Number(value=0.1, minimum=0.0, label="pairwise aggregate threshold") | |
| run_btn = gr.Button("推定を実行", variant="primary") | |
| with gr.Column(scale=3): | |
| with gr.Tab("データプレビュー"): | |
| data_shape = gr.Markdown( | |
| "### 実データサイズ: 未読み込み", | |
| elem_id="data-shape", | |
| ) | |
| data_preview = gr.Dataframe(label="Data preview", interactive=False, wrap=True) | |
| data_plot = gr.Plot(label="Data plot") | |
| with gr.Tab("結果サマリ"): | |
| summary = gr.Markdown() | |
| with gr.Tab("係数"): | |
| coef_fig = gr.Plot(label="Estimated coefficients") | |
| coef_df = gr.Dataframe(label="coefficients", interactive=False, wrap=True) | |
| with gr.Tab("グラフ構造"): | |
| graph_fig = gr.Plot(label="Estimated pairwise graph") | |
| pair_df = gr.Dataframe(label="pairwise summary", interactive=False, wrap=True) | |
| with gr.Tab("最適化ログ"): | |
| diag_fig = gr.Plot(label="Optimization diagnostics") | |
| with gr.Tab("エクスポート"): | |
| result_file = gr.File(label="Download results") | |
| gr.Markdown("---") | |
| gr.Markdown( | |
| """ | |
| # References | |
| - Original code: [kyanostat/mindemo](https://github.com/kyanostat/mindemo) | |
| - Original paper: [Sei and Yano, Minimum information dependence modeling, Bernoulli, 2024](https://projecteuclid.org/journals/bernoulli/volume-30/issue-4/Minimum-information-dependence-modeling/10.3150/23-BEJ1687.full) | |
| - Related work: [Sukeda and Sei, Minimum information Markov model, Journal of Multivariate Analysis, in press](https://www.sciencedirect.com/science/article/pii/S0047259X26000977) | |
| """ | |
| ) | |
| preview_btn.click(fn=preview_data, inputs=[csv_file, data_preset, data_mode], outputs=[data_preview, data_plot, data_shape]) | |
| run_btn.click( | |
| fn=run_experiment, | |
| inputs=[ | |
| csv_file, data_preset, data_mode, | |
| estimator, optimizer, lr, max_iter, tol, batch_size, stepsize_decay, | |
| l2_penalty, use_rpj, avg_start, | |
| fullbatch_l1, fullbatch_C, | |
| h_preset, custom_h_json, | |
| graph_threshold, random_seed, | |
| cle_L, cle_burnin, cle_thin, | |
| ], | |
| outputs=[summary, coef_fig, graph_fig, diag_fig, coef_df, pair_df, result_file], | |
| ) | |
| if __name__ == "__main__": | |
| server_name = os.environ.get("GRADIO_SERVER_NAME", "0.0.0.0") | |
| server_port = int(os.environ.get("GRADIO_SERVER_PORT", "7860")) | |
| demo.launch(server_name=server_name, server_port=server_port, css=CSS, show_error=True) |