"""计量课堂配套演示：供教师/AI 按段调用，返回模拟数据、OLS 和后验抽样结果。

只依赖 NumPy；导出教学图时另需 Matplotlib。导入不会读写文件；直接运行
会把三张图和数值摘要写入明确的输出目录，不读研究数据、不联网或安装包。
"""
from __future__ import annotations

import argparse
import json
from pathlib import Path

import numpy as np

SEED = 42
BLUE = "#3398E1"
INK = "#292929"
GREY = "#B7BEC4"


def generate_sample(n: int = 30, seed: int = SEED, *, heteroskedastic: bool = False):
    """返回 x/y/真实误差；两种生成方式均保持独立观测、零条件均值和已知斜率 2。"""
    if n < 3:
        raise ValueError("至少需要3个观察，才能估计截距/斜率并计算标准误。")
    rng = np.random.default_rng(seed)
    x = rng.uniform(-2, 2, n) if heteroskedastic else rng.normal(0, 1, n)
    sigma = 0.3 + 1.2 * np.abs(x) if heteroskedastic else np.ones(n)
    error = sigma * rng.normal(0, 1, n)
    return x, 1 + 2 * x + error, error


def fit_ols(x, y):
    """含截距的单变量OLS，返回系数、普通/HC1标准误、拟合值和残差。"""
    x, y = np.asarray(x, dtype=float), np.asarray(y, dtype=float)
    if x.ndim != 1 or y.shape != x.shape or len(x) <= 2:
        raise ValueError("x/y须为等长一维数组，且观察数大于2。")
    if not np.isfinite(x).all() or not np.isfinite(y).all():
        raise ValueError("示例不接受缺失或无穷值；先明确样本处理规则。")
    design = np.column_stack([np.ones(len(x)), x])
    beta, _, rank, _ = np.linalg.lstsq(design, y, rcond=None)
    if rank != 2:
        raise ValueError("解释变量必须有变化，否则无法识别斜率。")
    fitted = design @ beta
    residual = y - fitted
    n, k = design.shape
    bread = np.linalg.inv(design.T @ design)
    classical_cov = (residual @ residual) / (n - k) * bread
    # HC1只调整协方差估计，不重估系数；n/(n-k)是有限样本修正，k包括截距。
    scores = design * residual[:, None]
    hc1_cov = n / (n - k) * bread @ (scores.T @ scores) @ bread
    return {
        "beta": beta,
        "classical_se": np.sqrt(np.diag(classical_cov)),
        "hc1_se": np.sqrt(np.diag(hc1_cov)),
        "fitted": fitted,
        "residual": residual,
    }


def sampling_experiment(repetitions: int = 500, seed: int = SEED):
    """每次重新抽取独立样本；同一轮的30与300并非一个样本的前缀。"""
    if repetitions < 2:
        raise ValueError("重复抽样至少需要2次。")
    rng = np.random.default_rng(seed)
    slopes = {}
    for n in (30, 300):
        estimates = np.empty(repetitions)
        for i in range(repetitions):
            x = rng.normal(0, 1, n)
            y = 1 + 2 * x + rng.normal(0, 1, n)
            design = np.column_stack([np.ones(n), x])
            estimates[i] = np.linalg.lstsq(design, y, rcond=None)[0][1]
        slopes[n] = estimates
    return slopes


def make_figures(output_dir: Path):
    """仅由显式运行的main调用；写图后关闭画布，适合后台运行。"""
    import matplotlib
    matplotlib.use("Agg")
    import matplotlib.pyplot as plt

    plt.rcParams.update({"font.size": 11, "axes.labelcolor": INK, "text.color": INK})
    output_dir.mkdir(parents=True, exist_ok=True)
    x, y, _ = generate_sample()
    result = fit_ols(x, y)
    slopes = sampling_experiment()
    fig, axes = plt.subplots(1, 2, figsize=(11, 4.8), layout="constrained")
    fig.suptitle("OLS estimates across simulated samples", fontsize=17, y=1.08)
    axes[0].set_title("One sample: observations and fitted line", fontsize=11)
    axes[0].scatter(x, y, color=BLUE, edgecolor=INK, linewidth=0.4, s=30)
    grid = np.linspace(x.min(), x.max(), 100)
    axes[0].plot(grid, 1 + 2 * grid, color=INK, linestyle="--", label="True: y=1+2x")
    axes[0].plot(grid, result["beta"][0] + result["beta"][1] * grid, color=BLUE, label="Fitted OLS")
    axes[0].set(xlabel="x (simulated units)", ylabel="y (simulated units)")
    axes[0].legend(loc="upper left", frameon=False, fontsize=9)
    axes[1].set_title("Sampling distributions: n=30 versus n=300", fontsize=11)
    bins = np.linspace(min(a.min() for a in slopes.values()), max(a.max() for a in slopes.values()), 35)
    axes[1].hist(slopes[30], bins=bins, density=True, color=GREY, edgecolor=INK, linewidth=0.5, label="n=30")
    axes[1].hist(slopes[300], bins=bins, density=True, histtype="step", color=BLUE, linewidth=2, label="n=300")
    axes[1].axvline(2, color=INK, linestyle="--", linewidth=1, label="True slope=2")
    axes[1].set(xlabel="Estimated slope", ylabel="Density")
    axes[1].legend(frameon=False, fontsize=9)
    for ax in axes:
        ax.spines[["top", "right"]].set_visible(False)
        ax.grid(axis="y", color="#E8E8E8", linewidth=0.5)
        ax.set_axisbelow(True)
    fig.savefig(output_dir / "ols-sampling.png", dpi=160, bbox_inches="tight")
    plt.close(fig)

    hx, hy, _ = generate_sample(500, SEED, heteroskedastic=True)
    hetero = fit_ols(hx, hy)
    fig, axes = plt.subplots(1, 2, figsize=(11, 4.8), layout="constrained")
    fig.suptitle("Residuals and standard errors under heteroskedasticity", fontsize=16, y=1.08)
    axes[0].set_title("Residual spread varies with x", fontsize=11)
    axes[0].scatter(hx, hetero["residual"], color=BLUE, alpha=0.65, s=13)
    axes[0].axhline(0, color=INK, linestyle="--", linewidth=1)
    axes[0].set(xlabel="x (simulated units)", ylabel="OLS residual (simulated units)")
    axes[1].set_title("Same data and OLS coefficients; two SE estimates", fontsize=11)
    positions = np.arange(2)
    bars = [
        axes[1].bar(positions - 0.18, hetero["classical_se"], 0.36, color=BLUE, edgecolor=INK, linewidth=0.7, label="Conventional"),
        axes[1].bar(positions + 0.18, hetero["hc1_se"], 0.36, color="white", edgecolor=INK, linewidth=0.7, hatch="///", label="HC1"),
    ]
    for container in bars:
        axes[1].bar_label(container, fmt="%.3f", padding=3, fontsize=9)
    axes[1].set(xticks=positions, xticklabels=["Intercept", "Slope"], ylabel="Standard error (coefficient units)", ylim=(0, max(hetero["hc1_se"].max(), hetero["classical_se"].max()) * 1.4))
    axes[1].legend(frameon=False, loc="upper left", fontsize=9)
    for ax in axes:
        ax.spines[["top", "right"]].set_visible(False)
        ax.grid(axis="y", color="#E8E8E8", linewidth=0.5)
        ax.set_axisbelow(True)
    fig.savefig(output_dir / "heteroskedasticity.png", dpi=160, bbox_inches="tight")
    plt.close(fig)
    summary = {
        "seed": SEED,
        "simulated": True,
        "single_sample_beta": result["beta"].tolist(),
        "sampling": {str(n): {"repetitions": len(values), "mean_slope": float(values.mean()), "sd_slope": float(values.std(ddof=1))} for n, values in slopes.items()},
        "heteroskedasticity": {name: hetero[name].tolist() for name in ["beta", "classical_se", "hc1_se"]},
    }
    summary["mcmc"] = make_mcmc_figure(output_dir)
    (output_dir / "results.json").write_text(json.dumps(summary, indent=2) + "\n")
    return summary


def _bayes_inputs(x, y, prior_mean, prior_sd, sigma, intercept):
    """校验一维教学模型；允许任意有限 x，正规先验保证后验存在。"""
    x, y = np.asarray(x, dtype=float), np.asarray(y, dtype=float)
    if x.ndim != 1 or y.shape != x.shape or x.size == 0:
        raise ValueError("x/y 须为非空等长一维数组。")
    if not np.isfinite(x).all() or not np.isfinite(y).all():
        raise ValueError("教学模型不接受缺失或无穷值。")
    if not np.isfinite([prior_mean, prior_sd, sigma, intercept]).all() or prior_sd <= 0 or sigma <= 0:
        raise ValueError("先验标准差和已知误差标准差须为有限正数，均值和截距须有限。")
    return x, y


def beta_posterior(x, y, *, prior_mean=0.0, prior_sd=5.0, sigma=1.0, intercept=1.0):
    """固定截距/误差方差的正态回归斜率后验；解析结果仅作抽样核对参照。"""
    x, y = _bayes_inputs(x, y, prior_mean, prior_sd, sigma, intercept)
    variance = 1 / (1 / prior_sd**2 + (x @ x) / sigma**2)
    mean = variance * (prior_mean / prior_sd**2 + x @ (y - intercept) / sigma**2)
    return {"mean": float(mean), "sd": float(np.sqrt(variance))}


def metropolis_beta(x, y, *, step_sd=0.35, iterations=5000, seed=SEED,
                    initial=0.0, prior_mean=0.0, prior_sd=5.0, sigma=1.0, intercept=1.0):
    """用对称正态随机游走采样未归一化后验，返回每次迭代及全程接受率。"""
    x, y = _bayes_inputs(x, y, prior_mean, prior_sd, sigma, intercept)
    if not isinstance(iterations, int) or isinstance(iterations, bool) or iterations < 1:
        raise ValueError("iterations 须为正整数。")
    if not np.isfinite(step_sd) or step_sd <= 0 or not np.isfinite(initial):
        raise ValueError("proposal 标准差须为有限正数，初值须有限。")
    rng = np.random.default_rng(seed)

    def log_target(beta):
        residual = (y - intercept - beta * x) / sigma
        return -0.5 * (residual @ residual) - 0.5 * ((beta - prior_mean) / prior_sd)**2

    current = float(initial)
    current_logp = log_target(current)
    chain = np.empty(iterations)
    accepted = 0
    for i in range(iterations):
        proposal = current + rng.normal(0, step_sd)
        proposal_logp = log_target(proposal)
        # 对称 proposal 的 Hastings 比率为1；用 log 比较避免直接指数的数值溢出。
        if np.log(rng.random()) < min(0.0, proposal_logp - current_logp):
            current, current_logp = float(proposal), proposal_logp
            accepted += 1
        # 拒绝也必须保存当前状态；只留下接受点会改变目标分布及课堂解释。
        chain[i] = current
    return {"chain": chain, "acceptance_rate": accepted / iterations}


def make_mcmc_figure(output_dir: Path):
    """显式导出第三案例；固定同一数据比较步长，图与安全数值摘要不读外部文件。"""
    import matplotlib
    matplotlib.use("Agg")
    import matplotlib.pyplot as plt

    # 独立调用和完整 CLI 使用同一样式，不依赖前两个案例留下的绘图状态。
    plt.rcParams.update({"font.size": 11, "axes.labelcolor": INK, "text.color": INK})
    output_dir.mkdir(parents=True, exist_ok=True)
    x, y, _ = generate_sample()
    target = beta_posterior(x, y)
    steps = (0.01, 0.35, 3.0)
    runs = [metropolis_beta(x, y, step_sd=step) for step in steps]
    burn = 1000  # 教学比较统一剔除前1000次，不是对收敛的保证或自动调参。
    kept = [run["chain"][burn:] for run in runs]
    low = min(min(a.min() for a in kept), target["mean"] - 4 * target["sd"])
    high = max(max(a.max() for a in kept), target["mean"] + 4 * target["sd"])
    grid = np.linspace(low, high, 500)
    density = np.exp(-0.5 * ((grid - target["mean"]) / target["sd"])**2) / (target["sd"] * np.sqrt(2 * np.pi))
    bins = np.linspace(low, high, 40)
    peak = max(density.max(), *(np.histogram(draws, bins=bins, density=True)[0].max() for draws in kept))
    trace_low = min(0, *(run["chain"].min() for run in runs))
    trace_high = max(run["chain"].max() for run in runs)
    fig, axes = plt.subplots(3, 2, figsize=(11, 10), layout="constrained")
    try:
        fig.suptitle("Random-walk Metropolis: same posterior, different proposal scales", fontsize=15)
        for row, (step, run, draws) in enumerate(zip(steps, runs, kept)):
            axes[row, 0].plot(run["chain"], color=BLUE, linewidth=0.65)
            axes[row, 0].axvspan(0, burn, color=GREY, alpha=0.25)
            axes[row, 0].axhline(target["mean"], color=INK, linestyle="--", linewidth=1)
            axes[row, 0].set(title=f"Proposal SD={step:g}; acceptance={run['acceptance_rate']:.1%}", xlabel="Iteration", ylabel="Slope beta")
            axes[row, 1].hist(draws, bins=bins, density=True, color=GREY, alpha=0.8)
            axes[row, 1].plot(grid, density, color=BLUE, linewidth=2, label="Analytic posterior")
            axes[row, 1].set(title="After discarding first 1,000 iterations", xlabel="Slope beta", ylabel="Density")
            axes[row, 1].legend(frameon=False, fontsize=9)
            # 对照行使用共同坐标范围，避免自动缩放制造不同波动大小的错觉。
            axes[row, 0].set_ylim(trace_low - 0.1, trace_high + 0.1)
            axes[row, 1].set_ylim(0, peak * 1.1)
            for ax in axes[row]:
                ax.spines[["top", "right"]].set_visible(False)
                ax.grid(axis="y", color="#E8E8E8", linewidth=0.5)
                ax.set_axisbelow(True)
        fig.savefig(output_dir / "mcmc-posterior.png", dpi=160, bbox_inches="tight")
    finally:
        plt.close(fig)
    return {
        "data_seed": SEED, "chain_seed": SEED,
        "iterations": 5000, "discarded": burn, "analytic": target,
        "runs": {str(step): {"acceptance_rate": run["acceptance_rate"],
                              "mean": float(draws.mean()), "sd": float(draws.std(ddof=1)),
                              "credible_interval_95": np.quantile(draws, [0.025, 0.975]).tolist()}
                 for step, run, draws in zip(steps, runs, kept)},
    }


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description="Generate three simulated econometrics classroom examples.")
    parser.add_argument("--output-dir", type=Path, default=Path("classroom-results"))
    print(json.dumps(make_figures(parser.parse_args().output_dir), indent=2))
