MinDeMoDeMo / app.py
stardust-coder's picture
[add] first commit
d4afa71
Raw History Blame Contribute Delete
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)