"""PP5001 Week 5: run beside an empty folder called figures, or let this script create it."""
import pathlib as pl
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import scipy.stats as st
import statsmodels.formula.api as smf
import linearmodels.iv as iv

output = pl.Path('figures')
output.mkdir(exist_ok=True)

# These are schematic sampling distributions to distinguish two properties.
# They are normal curves chosen for illustration, not simulated IV densities.
grid = np.linspace(0, 10, 1500)
fig, axes = plt.subplots(1, 3, figsize=(12, 3.8), sharex=True, sharey=True)
for ax, title, means in zip(axes,
    ['Unbiased and consistent', 'Biased and consistent', 'Biased and inconsistent'],
    [[5, 5, 5], [6.5, 5.6, 5.15], [7, 7, 7]]):
    for mean, sd, label, color in zip(means, [1.2, .7, .35],
        ['Smaller n', 'Larger n', 'Much larger n'], ['#a9cce3', '#2980b9', '#154360']):
        ax.plot(grid, st.norm.pdf(grid, mean, sd), label=label, color=color)
    ax.axvline(5, color='#b03a2e', linestyle='--', label='True effect')
    ax.set_title(title, fontsize=12)
    ax.set_xlabel('Estimated effect')
    ax.set_yticks([])
axes[0].set_ylabel('Sampling density')
axes[1].legend(fontsize=9)
fig.tight_layout()
fig.savefig(output / 'unbiasedness-consistency.png', dpi=200)
plt.close(fig)

# The same population relationships as the original lecture illustration.
# w, e and v are independent, and Z is independent of all three.
# Seed 4 is an illustrative draw showing how IV can resemble OLS without relevance.
# Try other seeds: a zero first stage produces very unstable estimates.
def make_sample(n, first_stage_effect, seed=4):
    rng = np.random.default_rng(seed)
    z = (rng.random(n) < .4).astype(int)
    ability = rng.normal(0, 1, n)
    e = rng.normal(0, 10, n)
    v = rng.normal(0, 1, n)
    d = 1 + first_stage_effect * z + 2 * ability + v
    y = 2 + 5 * d + 20 * ability + e
    return pd.DataFrame({'y': y, 'x': d, 'z': z})

strong = make_sample(1000, 5)
weak_small = make_sample(1000, 0)
weak_large = make_sample(329509, 0)

# Fit each example explicitly. HC1 / robust request consistent standard errors.
# The usual IV standard errors can nevertheless be misleading without relevance.
strong_ols = smf.ols('y ~ x', data=strong).fit(cov_type='HC1')
strong_first = smf.ols('x ~ z', data=strong).fit(cov_type='HC1')
strong_iv = iv.IV2SLS.from_formula('y ~ 1 + [x ~ z]', data=strong).fit(cov_type='robust', debiased=True)
small_ols = smf.ols('y ~ x', data=weak_small).fit(cov_type='HC1')
small_first = smf.ols('x ~ z', data=weak_small).fit(cov_type='HC1')
small_iv = iv.IV2SLS.from_formula('y ~ 1 + [x ~ z]', data=weak_small).fit(cov_type='robust', debiased=True)
large_ols = smf.ols('y ~ x', data=weak_large).fit(cov_type='HC1')
large_first = smf.ols('x ~ z', data=weak_large).fit(cov_type='HC1')
large_iv = iv.IV2SLS.from_formula('y ~ 1 + [x ~ z]', data=weak_large).fit(cov_type='robust', debiased=True)

# Use the same axis limits in all three illustrations.
all_data = pd.concat([strong, weak_small, weak_large])
x_limits = (all_data['x'].min() - 1, all_data['x'].max() + 1)
y_limits = (all_data['y'].min() - 10, all_data['y'].max() + 10)

# Every observation is plotted, including all 329,509 in the large sample.
def draw_example(data, ols_model, iv_model, filename):
    means = data.groupby('z')[['x', 'y']].mean()
    fig, ax = plt.subplots(figsize=(10, 5.8))
    z0 = data.loc[data['z'] == 0]
    z1 = data.loc[data['z'] == 1]
    large = len(data) > 10000
    size, alpha = (2, .035) if large else (13, .4)
    ax.scatter(z0['x'], z0['y'], facecolors='none', edgecolors='#888888',
               s=size, alpha=alpha, linewidths=.5, label='z = 0', rasterized=True)
    ax.scatter(z1['x'], z1['y'], marker='^', color='#888888',
               s=size, alpha=alpha, label='z = 1', rasterized=True)
    grid = np.linspace(*x_limits, 200)
    ax.plot(grid, 2 + 5*grid, color='black', linewidth=2.8,
            label='True relationship: slope 5')
    ax.plot(grid, ols_model.params['Intercept'] + ols_model.params['x']*grid,
            color='#D55E00', linestyle='--', linewidth=2.8,
            label=f"OLS: {ols_model.params['x']:.2f}")
    ax.plot(grid, iv_model.params['Intercept'] + iv_model.params['x']*grid,
            color='#0072B2', linestyle='-.', linewidth=2.8,
            label=f"Wald / IV: {iv_model.params['x']:.2f}")
    ax.set_xlim(x_limits)
    ax.set_ylim(y_limits)
    # Project each group mean onto both axes. Offset labels so close means remain legible.
    for group in [0, 1]:
        mx, my = means.loc[group, ['x', 'y']]
        ax.plot([mx, mx], [y_limits[0], my], ':', color='black', linewidth=1.4)
        ax.plot([x_limits[0], mx], [my, my], ':', color='black', linewidth=1.4)
        offset = -22 if group == 0 else 22
        ax.annotate(r'$\bar{x}_{z=' + str(group) + '}$', xy=(mx, 0),
                    xycoords=('data', 'axes fraction'), xytext=(offset, -28),
                    textcoords='offset points', ha='center', fontsize=10,
                    arrowprops=dict(arrowstyle='-', color='black', lw=.6))
        ax.annotate(r'$\bar{y}_{z=' + str(group) + '}$', xy=(0, my),
                    xycoords=('axes fraction', 'data'), xytext=(-38, offset),
                    textcoords='offset points', ha='right', fontsize=10,
                    arrowprops=dict(arrowstyle='-', color='black', lw=.6))
    ax.scatter(means['x'], means['y'], marker='D', s=70, color='white',
               edgecolor='black', zorder=5, label='Instrument-group means')
    ax.set_xlabel('x', labelpad=32)
    ax.set_ylabel('y', labelpad=48)
    ax.set_title(f'{len(data):,} observations')
    ax.legend(fontsize=9, loc='upper left', framealpha=.95)
    fig.tight_layout()
    fig.savefig(output / filename, dpi=250)
    plt.close(fig)

draw_example(strong, strong_ols, strong_iv, 'wald-strong.png')
draw_example(weak_small, small_ols, small_iv, 'wald-zero-small.png')
draw_example(weak_large, large_ols, large_iv, 'wald-zero-large.png')

results = pd.DataFrame({
    'Example': ['Strong', 'No first stage, small', 'No first stage, large'],
    'N': [len(strong), len(weak_small), len(weak_large)],
    'OLS': [strong_ols.params['x'], small_ols.params['x'], large_ols.params['x']],
    'IV': [strong_iv.params['x'], small_iv.params['x'], large_iv.params['x']],
    'IV SE': [strong_iv.std_errors['x'], small_iv.std_errors['x'], large_iv.std_errors['x']],
    'First-stage coefficient': [strong_first.params['z'], small_first.params['z'], large_first.params['z']],
    'First-stage F': [float(strong_first.f_test('z=0').fvalue),
                      float(small_first.f_test('z=0').fvalue),
                      float(large_first.f_test('z=0').fvalue)],
    'First-stage R2': [strong_first.rsquared, small_first.rsquared, large_first.rsquared],
})
print(results.round(5).to_string(index=False))
results.to_csv(output / 'simulation-results.csv', index=False)
