{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "0f07d54c",
   "metadata": {},
   "source": [
    "# 실습 4 · 인과추론 실습 — DiD · 합성통제 · 반사실 예측 · DML 을 실데이터로\n",
    "\n",
    "한국외대 GBT 대학원 딥러닝 세미나 실습 4. 강의 페이지: `04_causal.html`. 인과추론 이론 시리즈(1~11장)의 실습편이다.\n",
    "\n",
    "순서: ① 설치 → ② 데이터 4종 로드 → ③ DiD → ④ 합성통제 → ⑤ 위약검정 → ⑥ 딥러닝 반사실(지하철) → ⑦ 기후동행카드 → ⑧ DoWhy 3단계 → ⑨ 반박 → ⑩ 관측자료 버전 → ⑪ CATE·군집 → ⑫ 보고표 → ⑬ 자기 데이터로 바꾸기\n",
    "\n",
    "Colab T4/CPU 기준 전체 약 10~15분(MLP seed 3회). 페이지의 숫자는 seed 5회·RTX 4090 실행값이라 소수점에서 조금 다를 수 있다.\n",
    "\n",
    "**Colab 사용법**: 이 파일을 다운로드해 [colab.research.google.com](https://colab.research.google.com) → 파일 → 노트 업로드."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "420f13c5",
   "metadata": {},
   "source": [
    "## ① 설치"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0f3a2579",
   "metadata": {},
   "outputs": [],
   "source": [
    "!pip -q install dowhy econml causaldata holidays"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "95d37a19",
   "metadata": {},
   "source": [
    "## ② 데이터 4종 로드\n",
    "한 줄 목표: 네 데이터의 크기와 열 이름을 확인한다. 관찰 포인트: 각 데이터에서 \"처치 단위·처치 시점·결과변수\"가 무엇인지 말할 수 있는가."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9a5384a6",
   "metadata": {},
   "outputs": [],
   "source": [
    "# pip install dowhy econml causaldata holidays   (Colab: 첫 셀에서 설치)\n",
    "import numpy as np, pandas as pd\n",
    "from causaldata import organ_donations\n",
    "import dowhy.datasets\n",
    "\n",
    "od  = organ_donations.load_pandas().data          # 27개 주 × 6분기 (2010Q4~2012Q1), 장기기증 등록률\n",
    "p99 = pd.read_csv(\"https://raw.githubusercontent.com/synth-inference/synthdid/master/data/california_prop99.csv\", sep=\";\")\n",
    "sub = pd.read_csv(\"https://hufs-ai-lecture.pages.dev/practice/data/subway_daily_total.csv\", parse_dates=[\"date\"])\n",
    "la  = dowhy.datasets.lalonde_dataset()            # NSW 실험표본 445명 (Dehejia & Wahba 1999)\n",
    "print(od.shape, p99.shape, sub.shape, la.shape)\n",
    "# → (162, 4) (1209, 4) (4261, 4) (445, 12)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a17735b8",
   "metadata": {},
   "outputs": [],
   "source": [
    "print(od.head(3)); print(p99.head(3)); print(sub.head(3)); print(la.head(3))"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "b70866e9",
   "metadata": {},
   "source": [
    "## ③ DiD — organ_donations (2011년 7월 캘리포니아 active choice)\n",
    "한 줄 목표: 네 칸 평균과 회귀 상호작용항이 같은 값을 내고, 클러스터 SE 가 기본 SE 와 얼마나 다른지 본다. 관찰 포인트: event study 의 사전 계수가 0 근처인가."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a96a6982",
   "metadata": {},
   "outputs": [],
   "source": [
    "import statsmodels.formula.api as smf\n",
    "od[\"Treat\"] = (od.State == \"California\").astype(int)\n",
    "od[\"Post\"]  = (od.Quarter_Num >= 4).astype(int)            # 2011Q3 부터 처치\n",
    "g = od.groupby([\"Treat\", \"Post\"]).Rate.mean()\n",
    "print((g[1,1] - g[1,0]) - (g[0,1] - g[0,0]))               # → -0.0225 (네 칸 평균)\n",
    "m = smf.ols(\"Rate ~ Treat*Post\", data=od).fit(\n",
    "        cov_type=\"cluster\", cov_kwds={\"groups\": od.State})  # 주 클러스터 SE (27개 주)\n",
    "print(m.params[\"Treat:Post\"], m.bse[\"Treat:Post\"])         # → -0.0225, 0.0061\n",
    "# event study: 기준 분기 2011Q2 (Quarter_Num 3) 를 빼고 처치군×분기 더미\n",
    "for k in [1, 2, 4, 5, 6]:\n",
    "    od[f\"T_q{k}\"] = od.Treat * (od.Quarter_Num == k)\n",
    "es = smf.ols(\"Rate ~ C(State) + C(Quarter_Num) + T_q1 + T_q2 + T_q4 + T_q5 + T_q6\",\n",
    "             data=od).fit(cov_type=\"cluster\", cov_kwds={\"groups\": od.State})\n",
    "print(es.params.filter(like=\"T_q\").round(4))\n",
    "# → T_q1 -0.0029  T_q2 0.0063  T_q4 -0.0216  T_q5 -0.0203  T_q6 -0.0222"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ea62fbac",
   "metadata": {},
   "outputs": [],
   "source": [
    "import matplotlib.pyplot as plt\n",
    "ci = es.conf_int(); ks = [1, 2, 3, 4, 5, 6]\n",
    "co = [0 if k == 3 else es.params[f\"T_q{k}\"] for k in ks]\n",
    "er = [[0 if k == 3 else co[i] - ci.loc[f\"T_q{k}\", 0] for i, k in enumerate(ks)], [0 if k == 3 else ci.loc[f\"T_q{k}\", 1] - co[i] for i, k in enumerate(ks)]]\n",
    "plt.errorbar(ks, co, yerr=er, fmt=\"o\", color=\"#A6432E\", capsize=4); plt.axhline(0, color=\"gray\"); plt.axvline(3.5, ls=\":\", color=\"gray\")\n",
    "plt.xticks(ks, [\"Q4'10\", \"Q1'11\", \"Q2'11\", \"Q3'11\", \"Q4'11\", \"Q1'12\"]); plt.ylabel(\"Treat × quarter (ref Q2'11)\"); plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "878bcbcd",
   "metadata": {},
   "source": [
    "## ④ 합성통제 — Prop 99 (1989)\n",
    "한 줄 목표: 비음·합 1 가중치를 scipy 로 직접 풀고, 사전 RMSPE 와 가중치 표를 얻는다. 관찰 포인트: 상위 가중치 주가 Abadie 외(2010)의 유타·네바다·몬태나와 겹치는가."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c3a11af5",
   "metadata": {},
   "outputs": [],
   "source": [
    "from scipy.optimize import minimize\n",
    "Y = p99.pivot(index=\"Year\", columns=\"State\", values=\"PacksPerCapita\")   # 31년 × 39주\n",
    "pre = np.asarray(Y.index <= 1988)\n",
    "def sc_weights(y1, Y0):                     # 비음·합 1 가중치, 사전 RMSPE 최소화\n",
    "    J = Y0.shape[1]\n",
    "    obj = lambda w: np.mean((y1 - Y0 @ w) ** 2)\n",
    "    jac = lambda w: -2 * Y0.T @ (y1 - Y0 @ w) / len(y1)\n",
    "    r = minimize(obj, np.full(J, 1/J), jac=jac, method=\"SLSQP\", bounds=[(0, 1)]*J,\n",
    "                 constraints={\"type\": \"eq\", \"fun\": lambda w: w.sum() - 1},\n",
    "                 options={\"ftol\": 1e-14, \"maxiter\": 2000})\n",
    "    w = np.clip(r.x, 0, None); return w / w.sum()\n",
    "def synth(unit, donors):\n",
    "    w = sc_weights(Y[unit].values[pre], Y[donors].values[pre])\n",
    "    gap = Y[unit].values - Y[donors].values @ w              # 관측 − 합성\n",
    "    rmspe = lambda mask: np.sqrt(np.mean(gap[mask] ** 2))\n",
    "    return w, gap, rmspe(pre), rmspe(~pre)\n",
    "donors = [s for s in Y.columns if s != \"California\"]\n",
    "w, gap, pre_r, post_r = synth(\"California\", donors)\n",
    "print(pd.Series(w, donors).sort_values(ascending=False).head(5).round(3))\n",
    "# → Utah 0.394  Montana 0.232  Nevada 0.205  Connecticut 0.109  New Hampshire 0.045\n",
    "print(pre_r, post_r, gap[~pre].mean(), gap[-1])             # → 1.66, 20.6, -19.5, -26.6"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1f4eb985",
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.plot(Y.index, Y[\"California\"], color=\"#A6432E\", lw=2.2, label=\"California\")\n",
    "plt.plot(Y.index, Y[donors].values @ w, color=\"#8A9AA1\", ls=\"--\", lw=2, label=\"synthetic\")\n",
    "plt.axvline(1988.5, ls=\":\", color=\"gray\"); plt.ylabel(\"packs per capita\"); plt.legend(); plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "ce15364e",
   "metadata": {},
   "source": [
    "## ⑤ in-space 위약검정\n",
    "한 줄 목표: 38개 주 각각을 처치로 두고 같은 절차를 돌려 사후/사전 RMSPE 비율의 순위로 p 를 만든다. 관찰 포인트: 캘리포니아보다 비율이 큰 주는 어느 주이고 그 주의 사전 적합은 어떤가."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "dc66e9ae",
   "metadata": {},
   "outputs": [],
   "source": [
    "res = {}\n",
    "for s in Y.columns:                                   # 39개 주 각각을 \"처치\"로 두고 같은 절차\n",
    "    d = donors if s == \"California\" else [x for x in Y.columns if x not in (s, \"California\")]\n",
    "    ws, g, r1, r2 = synth(s, d)\n",
    "    res[s] = dict(gap=g, pre=r1, post=r2, ratio=r2 / r1)\n",
    "ratio = pd.Series({s: v[\"ratio\"] for s, v in res.items()}).sort_values(ascending=False)\n",
    "rank = list(ratio.index).index(\"California\") + 1\n",
    "print(ratio.head(4).round(2)); print(\"rank\", rank, \"p =\", round(rank / 39, 3))\n",
    "# → Missouri 23.92  Virginia 19.83  California 12.44  Georgia 9.06 / rank 3 p = 0.077\n",
    "import matplotlib.pyplot as plt\n",
    "for s, v in res.items():\n",
    "    plt.plot(Y.index, v[\"gap\"], color=\"#8A9AA1\", lw=.8, alpha=.5)\n",
    "plt.plot(Y.index, res[\"California\"][\"gap\"], color=\"#A6432E\", lw=2.5); plt.axvline(1988.5, ls=\":\", color=\"gray\")\n",
    "plt.ylabel(\"gap = observed - synthetic\"); plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "fcbe84cc",
   "metadata": {},
   "source": [
    "## ⑥ 딥러닝 반사실 — 서울 지하철 코로나(2020-02-23)\n",
    "한 줄 목표: 개입 전 데이터로만 학습한 예측기(릿지 기준선·MLP)로 \"코로나가 없었을\" 승차를 예측하고, 위약 시점 검정의 잔차로 등각 구간을 만든다. 관찰 포인트: 위약 연도(2019-02~2020-01)의 누적 잔차가 2020년 격차에 비해 얼마나 작은가.\n",
    "\n",
    "규칙 ① 학습 데이터는 개입 전만(조기종료 검증도 개입 전). 규칙 ② 점예측에 반드시 구간을 붙인다."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d8f730d5",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch, torch.nn as nn, holidays\n",
    "from sklearn.linear_model import Ridge\n",
    "kr = holidays.KR(years=range(2015, 2027)); t0 = sub.date.min()\n",
    "def feats(d):                                     # 요일 7 + 공휴일 1 + 연중 Fourier 3차 6 + 선형 추세 1 = 15\n",
    "    doy = d.dt.dayofyear.values / 365.25\n",
    "    X = [np.eye(7)[d.dt.dayofweek.values], np.array([x in kr for x in d.dt.date])[:, None]]\n",
    "    for k in (1, 2, 3): X += [np.sin(2*np.pi*k*doy)[:, None], np.cos(2*np.pi*k*doy)[:, None]]\n",
    "    return np.hstack(X + [((d - t0).dt.days.values / 365.25)[:, None]]).astype(float)\n",
    "def fit_mlp(Xtr, ytr, Xte, seed, epochs=600, patience=40):\n",
    "    torch.manual_seed(seed); mu, sd = Xtr.mean(0), Xtr.std(0) + 1e-8; ym, ys = ytr.mean(), ytr.std()\n",
    "    Xt, yt = torch.tensor((Xtr-mu)/sd).float(), torch.tensor((ytr-ym)/ys).float()\n",
    "    net = nn.Sequential(nn.Linear(15,128), nn.ReLU(), nn.Linear(128,128), nn.ReLU(), nn.Linear(128,1))\n",
    "    opt = torch.optim.AdamW(net.parameters(), lr=2e-3, weight_decay=1e-3)\n",
    "    nv = len(Xt) // 10; best, bad = (1e9, None), 0   # 조기종료 검증 = 학습구간 마지막 10% (모두 개입 전)\n",
    "    for ep in range(epochs):\n",
    "        for b in torch.randperm(len(Xt) - nv).split(256):\n",
    "            opt.zero_grad(); loss = ((net(Xt[b]).squeeze() - yt[b])**2).mean(); loss.backward(); opt.step()\n",
    "        vl = ((net(Xt[-nv:]).squeeze() - yt[-nv:])**2).mean().item()\n",
    "        if vl < best[0]: best, bad = (vl, {k: v.clone() for k, v in net.state_dict().items()}), 0\n",
    "        elif (bad := bad + 1) >= patience: break\n",
    "    net.load_state_dict(best[1])\n",
    "    return net(torch.tensor((Xte-mu)/sd).float()).squeeze().detach().numpy() * ys + ym"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "bec4612d",
   "metadata": {},
   "outputs": [],
   "source": [
    "def counterfactual(train_end, p0, p1, seeds=(0, 1, 2), data=sub):\n",
    "    tr, te = data[data.date <= train_end], data[(data.date >= p0) & (data.date <= p1)]\n",
    "    Xtr, Xte, y = feats(tr.date), feats(te.date), tr.board.values / 1e6          # 백만 명\n",
    "    out = te[[\"date\"]].assign(obs=te.board.values / 1e6, ridge=Ridge(1.0).fit(Xtr, y).predict(Xte))\n",
    "    out[\"mlp\"] = np.mean([fit_mlp(Xtr, y, Xte, s) for s in seeds], axis=0)\n",
    "    return out\n",
    "plc  = counterfactual(\"2019-01-31\", \"2019-02-01\", \"2020-01-31\")   # 위약: 2019-02-23 을 가짜 개입으로\n",
    "main = counterfactual(\"2020-01-31\", \"2020-02-01\", \"2021-12-31\")   # 본 분석: 2020-02-23 코로나 '심각'\n",
    "r = (plc.obs - plc.mlp).values                                     # 위약 연도 잔차 → 등각 구간의 재료\n",
    "lo, hi = np.quantile(np.convolve(r, np.ones(30), \"valid\"), [.05, .95])   # 30일 합 잔차의 5·95% 분위수\n",
    "mo = main.assign(gap=main.obs - main.mlp, ym=main.date.dt.to_period(\"M\")).groupby(\"ym\").agg(gap=(\"gap\", \"sum\"), n=(\"gap\", \"size\"))\n",
    "mo[\"lo\"], mo[\"hi\"] = mo.gap + lo * mo.n / 30, mo.gap + hi * mo.n / 30   # 월별 격차와 90% 구간\n",
    "y20 = main[main.date.dt.year == 2020]\n",
    "print(f\"2020년 2~12월 손실 {(y20.obs - y20.mlp).sum():.0f}백만 명 ({100*(y20.obs - y20.mlp).sum()/y20.mlp.sum():.1f}%)\")\n",
    "print(f\"위약 연도 누적 잔차 {r.sum():+.0f}백만 명 ({100*r.sum()/plc.obs.sum():+.1f}%)  30일합 90% 구간 [{lo:.1f}, {hi:.1f}]\")\n",
    "# → 2020년 손실 약 -797백만 명 (-31%) · 위약 연도 잔차 +38 (+1.4%) · 구간 [-0.7, 8.1]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fbf2bc14",
   "metadata": {},
   "outputs": [],
   "source": [
    "fig, ax = plt.subplots(2, 1, figsize=(10, 6))\n",
    "ax[0].plot(main.date, main.obs.rolling(7, center=True).mean(), color=\"#A6432E\", label=\"observed\")\n",
    "ax[0].plot(main.date, main.mlp.rolling(7, center=True).mean(), color=\"#8A9AA1\", ls=\"--\", label=\"counterfactual (MLP)\")\n",
    "ax[0].plot(main.date, main.ridge.rolling(7, center=True).mean(), color=\"#2FA3A9\", ls=\":\", label=\"counterfactual (ridge)\"); ax[0].legend(); ax[0].set_ylabel(\"million/day\")\n",
    "xs = mo.index.to_timestamp(); ax[1].bar(xs, mo.gap, width=22, color=\"#A6432E\"); ax[1].errorbar(xs, mo.gap, yerr=[mo.gap - mo.lo, mo.hi - mo.gap], fmt=\"none\", ecolor=\"#002B49\", capsize=2)\n",
    "ax[1].axhline(0, color=\"gray\"); ax[1].set_ylabel(\"monthly gap (million)\"); plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "bec80cdf",
   "metadata": {},
   "source": [
    "## ⑦ 같은 절차를 기후동행카드(2024-01-27)에\n",
    "한 줄 목표: 효과가 작거나 없을 수 있는 개입에 같은 절차를 적용해, 격차가 구간 안에 들어오면 \"검출 불가\"로 보고하는 연습. 관찰 포인트: 학습 구간이 2년뿐이고 회복기 추세가 들어 있어 위약 잔차(=구간)가 코로나 분석보다 훨씬 넓다."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d3a44952",
   "metadata": {},
   "outputs": [],
   "source": [
    "sub2 = sub[sub.date >= \"2022-01-01\"].copy(); t0 = sub2.date.min()      # 추세 원점을 2022-01-01 로\n",
    "plc2  = counterfactual(\"2022-12-31\", \"2023-02-01\", \"2023-12-31\", data=sub2)  # 위약: 2023-01-27 가짜 개입\n",
    "main2 = counterfactual(\"2023-12-31\", \"2024-02-01\", \"2024-12-31\", data=sub2)  # 기후동행카드 2024-01-27\n",
    "r2 = (plc2.obs - plc2.mlp).values\n",
    "lo2, hi2 = np.quantile(np.convolve(r2, np.ones(30), \"valid\"), [.05, .95])\n",
    "mo2 = main2.assign(gap=main2.obs - main2.mlp, ym=main2.date.dt.to_period(\"M\")).groupby(\"ym\").agg(gap=(\"gap\", \"sum\"), n=(\"gap\", \"size\"))\n",
    "mo2[\"lo\"], mo2[\"hi\"] = mo2.gap + lo2 * mo2.n / 30, mo2.gap + hi2 * mo2.n / 30\n",
    "print(mo2.round(1)); print(\"구간이 0 을 포함한 달:\", int(((mo2.lo <= 0) & (mo2.hi >= 0)).sum()), \"/\", len(mo2))\n",
    "# → 11 / 11 → 이 설계로는 검출 불가 (격차 -2.7%, 90% 구간 약 ±10%)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "f7a4f44c",
   "metadata": {},
   "source": [
    "## ⑧ DoWhy 3단계 — LaLonde (식별 → 추정)\n",
    "한 줄 목표: 그래프를 명시해 backdoor 식별식을 얻고, 선형회귀·성향점수 매칭·LinearDML 세 추정치를 나란히 놓는다. 관찰 포인트: 실험표본이라 셋 다 1,794 근처에 있어야 정상이다."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7bf190d2",
   "metadata": {},
   "outputs": [],
   "source": [
    "from dowhy import CausalModel\n",
    "from sklearn.ensemble import GradientBoostingRegressor, GradientBoostingClassifier\n",
    "from sklearn.pipeline import make_pipeline; from sklearn.preprocessing import StandardScaler; from sklearn.linear_model import LogisticRegression\n",
    "PS = {\"propensity_score_model\": make_pipeline(StandardScaler(), LogisticRegression(max_iter=2000))}   # 성향점수 모형을 명시(수렴 보장)\n",
    "COV = [\"age\", \"educ\", \"black\", \"hisp\", \"married\", \"nodegr\", \"re74\", \"re75\"]\n",
    "graph = \"digraph { treat -> re78; \" + \" \".join(f\"{c} -> treat; {c} -> re78;\" for c in COV) + \" }\"\n",
    "la[\"treat\"] = la.treat.astype(int)\n",
    "model = CausalModel(data=la, treatment=\"treat\", outcome=\"re78\", graph=graph)   # ① 식별: 그래프를 명시\n",
    "estimand = model.identify_effect()                                             #    backdoor 식별식 출력\n",
    "e_lr  = model.estimate_effect(estimand, method_name=\"backdoor.linear_regression\")               # ② 추정 3종\n",
    "e_psm = model.estimate_effect(estimand, method_name=\"backdoor.propensity_score_matching\", target_units=\"att\", method_params=PS)\n",
    "np.random.seed(0); e_dml = model.estimate_effect(estimand, method_name=\"backdoor.econml.dml.LinearDML\",\n",
    "          method_params={\"init_params\": {\"model_y\": GradientBoostingRegressor(random_state=0),\n",
    "                                         \"model_t\": GradientBoostingClassifier(random_state=0),\n",
    "                                         \"discrete_treatment\": True, \"cv\": 5, \"random_state\": 0},\n",
    "                         \"fit_params\": {}})\n",
    "print(round(e_lr.value), round(e_psm.value), round(e_dml.value))\n",
    "# → 1676  2708  약 2000 (DML 은 교차적합 겹 분할 난수로 ±30) · 실험표본의 단순 평균 차이 1,794 가 벤치마크"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4110ff9e",
   "metadata": {},
   "outputs": [],
   "source": [
    "print(estimand)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "fb7d60f3",
   "metadata": {},
   "source": [
    "## ⑨ 반박 3종\n",
    "한 줄 목표: placebo treatment · random common cause · data subset 이 모두 \"통과\"인지 표로 만든다. 관찰 포인트: 통과는 인과의 증명이 아니라 \"탈락하면 확실히 문제\"라는 비대칭 검정이다."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c58e045b",
   "metadata": {},
   "outputs": [],
   "source": [
    "refs = {}\n",
    "for name, kw in [(\"placebo_treatment_refuter\", {\"placebo_type\": \"permute\"}),   # 처치를 무작위로 섞기 → 0 근처여야 통과\n",
    "                 (\"random_common_cause\", {}),                                    # 가짜 교란 추가 → 추정치가 안 흔들려야\n",
    "                 (\"data_subset_refuter\", {\"subset_fraction\": 0.8})]:            # 80% 부분표본 → 안정적이어야\n",
    "    rf = model.refute_estimate(estimand, e_dml, method_name=name, num_simulations=20, random_seed=0, **kw)\n",
    "    refs[name] = (round(rf.estimated_effect), round(rf.new_effect), round(rf.refutation_result[\"p_value\"], 2))\n",
    "print(pd.DataFrame(refs, index=[\"원 추정치\", \"새 추정치\", \"p\"]).T)             # ③ 반박\n",
    "# → placebo 1995 → -119 (p .44) · random_cc 1995 → 1820 (p .24) · subset 1995 → 1775 (p .35)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "86264370",
   "metadata": {},
   "source": [
    "## ⑩ 관측자료 버전: NSW 처치군 + PSID 대조군\n",
    "한 줄 목표: 같은 코드에 대조군만 실험 대조군에서 PSID 조사 대조군으로 바꿔, 식별이 무너지면 ML 도 소용없음을 본다. 관찰 포인트: 단순 차이 -15,205 → 세 추정치가 1,794 를 얼마나 회복하는가."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "da38dd6e",
   "metadata": {},
   "outputs": [],
   "source": [
    "# 같은 파이프라인을 관측자료 버전(NSW 처치군 185 + PSID 대조군 2,490)에 그대로 적용\n",
    "nsw  = pd.read_stata(\"https://users.nber.org/~rdehejia/data/nsw_dw.dta\")\n",
    "psid = pd.read_stata(\"https://users.nber.org/~rdehejia/data/psid_controls.dta\")\n",
    "ob = pd.concat([nsw[nsw.treat == 1], psid]).rename(columns={\"education\": \"educ\", \"hispanic\": \"hisp\", \"nodegree\": \"nodegr\"})\n",
    "ob = ob.reset_index(drop=True).astype({c: float for c in COV + [\"re78\"]}); ob[\"treat\"] = ob.treat.astype(int)\n",
    "m2 = CausalModel(data=ob, treatment=\"treat\", outcome=\"re78\", graph=graph); es2 = m2.identify_effect()\n",
    "print(ob[ob.treat == 1].re78.mean() - ob[ob.treat == 0].re78.mean())                        # → -15205 (단순 차이)\n",
    "print(round(m2.estimate_effect(es2, method_name=\"backdoor.linear_regression\").value),        # → 752\n",
    "      round(m2.estimate_effect(es2, method_name=\"backdoor.propensity_score_matching\", target_units=\"att\", method_params=PS).value),  # → 2697\n",
    "      round(m2.estimate_effect(es2, method_name=\"backdoor.econml.dml.LinearDML\",\n",
    "              method_params={\"init_params\": {\"model_y\": GradientBoostingRegressor(random_state=0),\n",
    "                                             \"model_t\": GradientBoostingClassifier(random_state=0),\n",
    "                                             \"discrete_treatment\": True, \"cv\": 5, \"random_state\": 0}, \"fit_params\": {}}).value))  # → 약 -600"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "6f490c0f",
   "metadata": {},
   "source": [
    "## ⑪ CATE 와 하위집단 군집\n",
    "한 줄 목표: CausalForestDML 의 개인별 효과 분포를 그리고, 공변량 K-means(k=4)로 \"효과가 큰 하위집단\"의 프로파일 표를 만든다. 관찰 포인트: 군집 평균 CATE 의 90% 구간이 0 을 제외하는 군집이 있는가."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0d7fed69",
   "metadata": {},
   "outputs": [],
   "source": [
    "from econml.dml import CausalForestDML\n",
    "from sklearn.preprocessing import StandardScaler\n",
    "from sklearn.cluster import KMeans\n",
    "X = la[COV].values.astype(float)\n",
    "cf = CausalForestDML(model_y=GradientBoostingRegressor(random_state=0), model_t=GradientBoostingClassifier(random_state=0),\n",
    "                     discrete_treatment=True, n_estimators=1000, min_samples_leaf=10, cv=5, random_state=0)\n",
    "cf.fit(la.re78.values, la.treat.values, X=X)\n",
    "cate = cf.effect(X); lo_i, hi_i = cf.effect_interval(X, alpha=.1)\n",
    "print(f\"ATE {cf.ate(X):.0f}  CATE sd {cate.std():.0f}  90% 구간이 0 을 제외하는 비율 {np.mean((lo_i > 0) | (hi_i < 0)):.2f}\")\n",
    "# → ATE 1948  sd 1051  비율 0.49\n",
    "lab = KMeans(4, n_init=20, random_state=0).fit_predict(StandardScaler().fit_transform(X))   # 공변량 군집 k=4\n",
    "rows = []\n",
    "for k in range(4):\n",
    "    ai = cf.ate_inference(X[lab == k]); ci = ai.conf_int_mean(alpha=.1)\n",
    "    rows.append({\"군집\": k, \"n\": int((lab == k).sum()), **la[COV][lab == k].mean().round(2).to_dict(),\n",
    "                 \"CATE\": round(float(ai.mean_point)), \"90%하\": round(float(ci[0])), \"90%상\": round(float(ci[1]))})\n",
    "print(pd.DataFrame(rows).sort_values(\"CATE\", ascending=False).to_string(index=False))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f096c88b",
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.hist(cate, bins=30, color=\"#146E7A\"); plt.axvline(cf.ate(X), color=\"#A6432E\"); plt.xlabel(\"CATE (1978 earnings, $)\"); plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "fb3ef7bb",
   "metadata": {},
   "source": [
    "## ⑫ 보고표 생성\n",
    "11장 체크리스트(식별 가정 → 정황 증거 → 불확실성 → 기준선 → 벤치마크 한계)를 열로 갖는 표를 코드로 만든다."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7a8879cb",
   "metadata": {},
   "outputs": [],
   "source": [
    "# 11장 체크리스트를 결과표로 채운다 — 논문 부록에 그대로 붙일 수 있는 형태\n",
    "report = pd.DataFrame([\n",
    "  [\"DiD (organ)\",   \"평행추세·무예견\",  \"event study 사전 계수 (Q4'10 -0.003, Q1'11 +0.006)\", f\"{m.params['Treat:Post']:.4f} (클러스터 SE {m.bse['Treat:Post']:.4f}, 27개 주)\", \"TWFE 동일 계수\"],\n",
    "  [\"합성통제 (Prop99)\",\"볼록결합·사전 적합\", f\"사전 RMSPE {pre_r:.2f}, 가중치 6개 주 > 0.01\", f\"사후 평균 격차 {gap[~pre].mean():.1f}갑, 2000년 {gap[-1]:.1f}갑\", f\"in-space 위약 순위 {rank}/39 (p={rank/39:.3f})\"],\n",
    "  [\"딥러닝 반사실 (지하철)\",\"학습=개입 전만\",  f\"위약 연도 잔차 {100*r.sum()/plc.obs.sum():+.1f}%\", f\"2020년 손실 {(y20.obs-y20.mlp).sum():.0f}백만 명\", \"등각 90% 구간(위약 잔차)\"],\n",
    "  [\"DML (LaLonde)\", \"무교란(backdoor)\", \"estimand 출력 부록 수록\", f\"OLS {e_lr.value:.0f} / PSM {e_psm.value:.0f} / DML {e_dml.value:.0f}\", \"반박 3종 통과 (표)\"],\n",
    "], columns=[\"분석\", \"식별 가정\", \"정황 증거\", \"추정치·불확실성\", \"반박·기준선\"])\n",
    "print(report.to_string(index=False))"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "eee39c76",
   "metadata": {},
   "source": [
    "## ⑬ 자기 데이터로 바꾸기\n",
    "\n",
    "| 내 데이터의 모양 | 바꿀 셀 | 이 노트북의 자리 |\n",
    "|---|---|---|\n",
    "| 여러 단위 × 몇 시점 패널, 일부 단위만 처치 | ② 의 `od` 를 내 패널로, `Treat`·`Post` 정의만 수정 | ③ DiD |\n",
    "| 처치 단위 1개 + 대조 단위 여러 개의 긴 시계열 | ② 의 `p99` 를 (Year, State, Y) 긴 형식으로 | ④⑤ 합성통제 |\n",
    "| 처치 단위 1개의 긴 일별·주별 시계열 (대조 없음) | ② 의 `sub` 를 (date, board) 로, 개입일과 `feats` 의 캘린더 특징만 수정 | ⑥⑦ 반사실 예측 |\n",
    "| 개인 횡단면, 처치 더미 + 공변량 | ② 의 `la` 를 내 표로, `COV` 와 `graph` 만 수정 | ⑧~⑪ DoWhy·DML |\n",
    "\n",
    "바꾸지 말아야 할 것: 학습 데이터의 경계(개입 전만), 클러스터 수준(처치가 배정된 단위), 반박 3종. 바꿔야 할 것: 식별 가정을 내 설계의 말로 다시 적기."
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "name": "python"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
