# homework4_solution.jl
# Computational Bootcamp, Summer 2026 -- Homework 4 solutions
#
# Q1  Aiyagari: EGM, two distribution methods, and two market-clearing methods
# Q2  A multinomial logit: MLE, standard errors, and an IIA counterfactual
# Q3  SMM on the Question 1 model
#
#     julia --project=. homework4.jl          # ~15 minutes end to end

using CSV, DataFrames, ForwardDiff, Interpolations, LinearAlgebra
using Optim, Plots, Random, Roots, Statistics

DATA_DIR = joinpath(@__DIR__, "data")
FIGDIR   = joinpath(@__DIR__, "figures")
mkpath(FIGDIR)

rd(x, d = 4) = round(x, digits = d)    # rounding helpers for printing
sg(x, d = 3) = round(x, sigdigits = d)

# =====================================================================
# Question 1: An Aiyagari economy
# =====================================================================

# --- Part A: the firm block ------------------------------------------------
# r = α K^(α-1) - δ and w = (1-α)K^α with L = 1; inverting the first gives
# capital DEMAND, which must slope DOWN in r.
capital_demand(r, para) = (para.α / (r + para.δ))^(1 / (1 - para.α))
wage(r, para) = (1 - para.α) * capital_demand(r, para)^para.α

# Complete-markets benchmark.
K_complete_markets(para) = capital_demand(1/para.β - 1, para)

@kwdef struct AiyagariParameters
    γ::Float64 = 2.0                       # CRRA
    β::Float64 = 0.96
    α::Float64 = 0.36
    δ::Float64 = 0.08

    # Changing p_ue updates unemployment and Π while preserving E[y] = 1.
    p_eu::Float64 = 0.05                                  # separation rate
    p_ue::Float64 = 0.75                                  # job-finding rate
    u_rate::Float64 = p_eu / (p_eu + p_ue)                # ergodic unemployment
    y_e::Float64 = 1 / (1 - u_rate)                       # so E[y] = L = 1
    y_grid::Vector{Float64} = [y_e, 0.0]                  # employed, unemployed
    Π::Matrix{Float64} = [1-p_eu p_eu; p_ue 1-p_ue]       # Π[i,j] = P(y'=j | y=i)
    N_y::Int64 = length(y_grid)

    # Zero is the natural borrowing limit; concentrate grid points near it.
    a_min::Float64 = 0.0
    a_max::Float64 = 50.0
    N_a::Int64 = 300
    curve::Float64 = 2.0
    a_grid::Vector{Float64} =
        a_min .+ (a_max - a_min) .* (collect(range(0, 1; length = N_a)) .^ curve)

    N_sim::Int64 = 10_000                  # households in the simulated panel
    T_sim::Int64 = 500                     # dates (burn-in checked in part B)
    seed::Int64 = 20260826

    # Floor consumption to avoid u'(0) = Inf at (a, y) = (0, 0).
    c_min::Float64 = 1e-10

    tol::Float64 = 1e-9                    # EGM, on the consumption policy
    max_iter::Int64 = 10_000
    tol_dist::Float64 = 1e-11              # the distribution is a LINEAR fixed
    max_iter_dist::Int64 = 200_000         # point, so ask for more digits
end

# Mutable: the solvers write into it, and part C's outer loops move r and w.
mutable struct AiyagariSolutions
    c_pol::Matrix{Float64}                 # N_a × N_y: c(a, y)
    a_pol::Matrix{Float64}                 # N_a × N_y: a'(a, y)
    μ::Matrix{Float64}                     # N_a × N_y: the histogram
    T_star::Matrix{Float64}                # (N_a N_y)²: the Young transition
    a_sim::Vector{Float64}                 # N_sim: the simulated cross-section
    y_sim::Vector{Int64}
    r::Float64                             # prices, read by the operators
    w::Float64
    K::Float64                             # aggregates, set by aggregate_savings
    Y::Float64
end

function initialize(; r = 0.03, kwargs...)
    para = AiyagariParameters(; kwargs...)
    w = wage(r, para)
    c_pol = [max((1 + r)*para.a_grid[i] + w*para.y_grid[iy], para.c_min)     # guess:
             for i in 1:para.N_a, iy in 1:para.N_y]                         # consume all
    sols = AiyagariSolutions(c_pol, zeros(para.N_a, para.N_y),
                             ones(para.N_a, para.N_y) ./ (para.N_a * para.N_y),
                             zeros(para.N_a*para.N_y, para.N_a*para.N_y),
                             zeros(para.N_sim), ones(Int64, para.N_sim),
                             r, w, 0.0, 0.0)
    return para, sols
end

u_prime(c, γ) = c^(-γ)
u_prime_inv(x, γ) = x^(-1/γ)               # closed form, and the point of EGM

# --- Part A: the EGM operator ----------------------------------------------
# HW3's EGM steps with labor income w*y and a consumption floor.
function egm_step(para, sols)
    (; γ, β, y_grid, Π, a_grid, a_min, N_a, N_y, c_min) = para
    (; c_pol, r, w) = sols

    cf = [linear_interpolation(a_grid, c_pol[:, iy]; extrapolation_bc = Line())
          for iy in 1:N_y]
    c_next = similar(c_pol)
    RHS = zeros(N_a)

    for iy in 1:N_y
        # 1. Euler RHS at each point of TOMORROW's asset grid, then invert u'.
        for j in 1:N_a
            s = 0.0
            for jp in 1:N_y
                s += Π[iy, jp] * u_prime(max(cf[jp](a_grid[j]), c_min), γ)
            end
            RHS[j] = β * (1 + r) * s
        end
        c_endog = u_prime_inv.(RHS, γ)

        # 2. The budget constraint is linear in a, so today's assets fall out:
        #    a = (c + a' - w y)/(1 + r). This endogenous grid differs across y.
        a_endog = (c_endog .+ a_grid .- w*y_grid[iy]) ./ (1 + r)

        # 3. Interpolate back onto the fixed grid; below the smallest
        #    endogenous asset level (step 4) the constraint binds, so set
        #    a' = a_min and read c off the budget line.
        c_of_a = linear_interpolation(a_endog, c_endog; extrapolation_bc = Line())
        for i in 1:N_a
            c_next[i, iy] = a_grid[i] <= a_endog[1] ?
                max((1 + r)*a_grid[i] + w*y_grid[iy] - a_min, c_min) :
                c_of_a(a_grid[i])
        end
    end
    return c_next
end

function solve_egm!(para, sols)
    (; tol, max_iter, a_grid, y_grid, a_min, a_max, N_a, N_y) = para
    max_diff, n = tol + 1.0, 0
    while max_diff > tol && n < max_iter
        n += 1
        c_next = egm_step(para, sols)
        max_diff = maximum(abs.(c_next .- sols.c_pol))
        sols.c_pol .= c_next
    end
    n == max_iter && @warn "EGM did not converge in $max_iter iterations"
    for iy in 1:N_y, i in 1:N_a            # a' from the budget constraint
        sols.a_pol[i, iy] = clamp((1 + sols.r)*a_grid[i] + sols.w*y_grid[iy] -
                                  sols.c_pol[i, iy], a_min, a_max)
    end
    return n
end

# Update both prices before solving the household problem.
function set_prices!(para, sols, r)
    sols.r, sols.w = r, wage(r, para)
    return sols
end

# --- Part B, method 1: 10,000 households on the interpolated policy ---------
# Fix uniforms, not income states, so the same draws work when Π changes.
draw_shocks(para) = rand(Xoshiro(para.seed), para.N_sim, para.T_sim)

function markov_step(j, Π, x)
    cum = 0.0
    for k in axes(Π, 2)
        cum += Π[j, k]
        x <= cum && return k
    end
    return size(Π, 2)
end

function simulate_y_panel(para; u = nothing)
    (; Π, N_sim, T_sim, u_rate) = para
    u === nothing && (u = draw_shocks(para))
    y = ones(Int64, N_sim, T_sim)
    for i in 1:N_sim
        y[i, 1] = u[i, 1] < u_rate ? 2 : 1     # start from the chain's own
        for t in 2:T_sim                       # stationary distribution
            y[i, t] = markov_step(y[i, t-1], Π, u[i, t])
        end
    end
    return y
end

# Simulated households are not confined to the grid, so interpolate the policy.
a_pol_interp(para, sols) =
    [linear_interpolation(para.a_grid, sols.a_pol[:, iy]; extrapolation_bc = Line())
     for iy in 1:para.N_y]

function simulate_panel!(para, sols, y_panel; a0 = 0.0)
    (; a_min, a_max, N_sim) = para
    T = size(y_panel, 2)
    g = a_pol_interp(para, sols)
    a = fill(a0, N_sim)
    path = zeros(T)                        # cross-sectional mean by date
    path[1] = mean(a)
    for t in 2:T
        @inbounds for i in 1:N_sim
            a[i] = clamp(g[y_panel[i, t-1]](a[i]), a_min, a_max)
        end
        path[t] = mean(a)
    end
    sols.a_sim .= a
    sols.y_sim .= y_panel[:, T]
    return sols.a_sim, path
end

# --- Part B, method 2: the histogram ---------------------------------------
# Split mass between bracketing nodes, preserving mean assets a'.
function lottery(ap, grid)
    ap = clamp(ap, first(grid), last(grid))
    k = min(searchsortedlast(grid, ap), length(grid) - 1)
    return k, k + 1, (grid[k+1] - ap) / (grid[k+1] - grid[k])
end

flat(i, j, N_a) = (j - 1)*N_a + i

function transition_matrix(para, sols)
    (; Π, a_grid, N_a, N_y) = para
    T_star = zeros(N_a*N_y, N_a*N_y)
    for j in 1:N_y, i in 1:N_a
        lo, hi, ω = lottery(sols.a_pol[i, j], a_grid)
        for jp in 1:N_y                    # a' and y' independent given today
            T_star[flat(i,j,N_a), flat(lo,jp,N_a)] += ω * Π[j, jp]
            T_star[flat(i,j,N_a), flat(hi,jp,N_a)] += (1 - ω) * Π[j, jp]
        end
    end
    return T_star
end

# Iterate μ' = T*'μ, retaining μ as a warm start for subsequent solves.
function stationary_distribution!(para, sols)
    (; tol_dist, max_iter_dist) = para
    sols.T_star .= transition_matrix(para, sols)
    Tt = transpose(sols.T_star)

    μ = vec(sols.μ)                        # a view: writing μ writes sols.μ
    μ ./= sum(μ)
    μ_next = similar(μ)
    max_diff, n = tol_dist + 1.0, 0
    while max_diff > tol_dist && n < max_iter_dist
        n += 1
        mul!(μ_next, Tt, μ)
        max_diff = maximum(abs.(μ_next .- μ))
        copyto!(μ, μ_next)
    end
    n == max_iter_dist && @warn "stationary distribution did not converge"
    μ ./= sum(μ)
    return n
end

# --- Part B: wealth shares -------------------------------------------------
# Wealth held by the richest fraction p, using fractional mass at the cutoff.
# Works with either histogram weights or equally weighted simulated households.
function top_wealth_share(a, μ, p)
    a = collect(a)
    m = collect(μ) ./ sum(μ)
    idx = sortperm(a; rev = true)          # richest first
    a, m = a[idx], m[idx]

    total = sum(m .* a)
    total <= 0 && return 0.0
    remaining, wealth = float(p), 0.0
    for k in eachindex(a)
        remaining <= 0 && break
        take = min(m[k], remaining)
        wealth += take * a[k]
        remaining -= take
    end
    return wealth / total
end

wealth_distribution(para, sols; method = :histogram) =
    method === :histogram ? (para.a_grid, vec(sum(sols.μ, dims = 2))) :
                            (sols.a_sim, fill(1/length(sols.a_sim), length(sols.a_sim)))

# --- Part B/C: aggregate savings and market clearing -----------------------
# A(r): solve the household problem at (r, w(r)), find the stationary
# distribution, integrate assets. It WRITES into sols, so callers warm-start.
function aggregate_savings(r, para, sols; method = :histogram,
                           y_panel = nothing, shocks = nothing)
    set_prices!(para, sols, r)
    solve_egm!(para, sols)
    if method === :histogram
        stationary_distribution!(para, sols)
    else
        simulate_panel!(para, sols,
                        y_panel === nothing ? simulate_y_panel(para; u = shocks) : y_panel)
    end
    a, m = wealth_distribution(para, sols; method = method)
    sols.K = sum(m .* a)
    sols.Y = sols.K^para.α                 # L = 1
    return sols.K
end

excess_demand(r, para, sols; kwargs...) =
    aggregate_savings(r, para, sols; kwargs...) - capital_demand(r, para)

# Bracket equilibrium below 1/β - 1. Try a narrower bracket near r_guess.
function clear_market!(para, sols; method = :histogram, shocks = nothing,
                       r_lo = -0.01, r_hi = nothing, r_guess = nothing, xatol = 1e-6)
    r_hi === nothing && (r_hi = 1/para.β - 1)
    # Hold the income panel fixed throughout the root search.
    yp = method === :panel ? simulate_y_panel(para; u = shocks) : nothing
    f(r) = excess_demand(r, para, sols; method = method, y_panel = yp)

    lo, hi = r_lo, r_hi
    if r_guess !== nothing                 # try a narrow bracket first
        step = 0.002
        for _ in 1:5
            lo, hi = max(r_guess - step, r_lo), min(r_guess + step, r_hi)
            f(lo)*f(hi) < 0 && break
            step *= 3
            (lo == r_lo && hi == r_hi) && break
        end
        f(lo)*f(hi) < 0 || ((lo, hi) = (r_lo, r_hi))
    end

    r_star = find_zero(f, (lo, hi), Roots.Brent(); xatol = xatol)
    K_star = aggregate_savings(r_star, para, sols; method = method, y_panel = yp)
    return r_star, K_star, para, sols
end

# Equilibrium by the damped update K' = (1-λ)K + λ A(r(K)). See q1() for λ.
function clear_market_damped!(para, sols; λ = 0.5, method = :histogram,
                              shocks = nothing, K0 = nothing, tol = 1e-6,
                              max_iter = 500)
    (; α, δ, β) = para
    yp = method === :panel ? simulate_y_panel(para; u = shocks) : nothing
    # The firm's FOC read the other way, capped below 1/β - 1: there the
    # household problem has no stationary solution.
    r_safe(K) = min(α*K^(α - 1) - δ, 1/β - 1 - 1e-6)

    K = K0 === nothing ? K_complete_markets(para) : K0
    path = Float64[K]
    max_diff, n = tol + 1.0, 0
    while max_diff > tol && n < max_iter
        n += 1
        A = aggregate_savings(r_safe(K), para, sols; method = method, y_panel = yp)
        K_new = (1 - λ)*K + λ*A
        max_diff = abs(K_new - K)
        K = K_new
        push!(path, K)
    end
    converged = max_diff <= tol
    converged || @warn "damped iteration did not converge in $max_iter steps (λ = $λ)"
    # Refresh at the returned rate; sols.K is supplied savings, which equals
    # the capital iterate K only at equilibrium.
    r = r_safe(K)
    aggregate_savings(r, para, sols; method = method, y_panel = yp)
    return r, K, para, sols, n, converged, path
end

function q1()
    para, sols = initialize()
    @assert all(sum(para.Π, dims = 2) .≈ 1.0)
    @assert para.y_grid' * [1 - para.u_rate, para.u_rate] ≈ 1.0   # E[y] = L = 1
    row(lbl, h, p) = println("       ", rpad(lbl, 20), lpad(h, 11), lpad(p, 13))

    # --- Part A -----------------------------------------------------------
    r0 = 0.03
    set_prices!(para, sols, r0)
    t0 = time(); n_egm = solve_egm!(para, sols); t_egm = time() - t0
    println("Q1A: u = ", rd(para.u_rate), "  y_e = ", rd(para.y_e), "  w(", r0,
            ") = ", rd(sols.w, 5), "  K^d(", r0, ") = ", rd(capital_demand(r0, para)))
    println("     EGM: ", n_egm, " iterations, ", rd(t_egm, 2), " s;  max a' = ",
            rd(maximum(sols.a_pol), 3), " (a_max = ", para.a_max, ")")
    @assert maximum(sols.a_pol) < para.a_max        # the grid must not be the model

    pp = plot(para.a_grid, sols.a_pol; lw = 2, label = ["employed" "unemployed"],
              xlabel = "a", ylabel = "a'", xlims = (0, 20),
              title = "Savings policy at r = $r0")
    plot!(pp, para.a_grid, para.a_grid; ls = :dot, c = :black, label = "45°")
    savefig(pp, joinpath(FIGDIR, "hw4_q1_policy.png"))
    pc = plot(para.a_grid, sols.c_pol; lw = 2, label = ["employed" "unemployed"],
              xlabel = "a", ylabel = "c", xlims = (0, 20),
              title = "Consumption policy at r = $r0")
    savefig(pc, joinpath(FIGDIR, "hw4_q1_consumption.png"))
    # c(0, unemployed) = 0 and rises very steeply just above: a household with
    # nothing and no income is a hair from zero, so precaution is enormous.

    # --- Part B -----------------------------------------------------------
    y_panel = simulate_y_panel(para)
    println("Q1B: simulated unemployment rate = ", rd(mean(y_panel[:, end] .== 2)),
            " (ergodic ", rd(para.u_rate), ")")

    # Compare initial assets 0 and 12 under identical shocks to measure burn-in.
    _, path0 = simulate_panel!(para, sols, y_panel; a0 = 0.0)
    a_mc = copy(sols.a_sim)
    _, path1 = simulate_panel!(para, sols, y_panel; a0 = 12.0)
    simulate_panel!(para, sols, y_panel; a0 = 0.0)      # leave sols at a0 = 0
    println("     burn-in: |mean assets from a0 = 0  −  from a0 = 12| stays below")
    for tolm in (1e-2, 1e-6, 1e-10)
        j = something(findfirst(t -> all(abs.(path0[t:end] .- path1[t:end]) .< tolm),
                                1:para.T_sim), para.T_sim)
        println("       ", tolm, " from t = ", j, " on")
    end
    println("       cross-section taken at T = ", para.T_sim)
    # Under 0.3% of mean assets by t ≈ 136 and 1e-6 by t ≈ 290, so T = 500 has
    # a wide margin.
    pb = plot(1:para.T_sim, [path0 path1]; lw = 2, ls = [:solid :dash],
              label = ["start a0 = 0" "start a0 = 12"], xlabel = "date",
              ylabel = "cross-sectional mean assets", title = "Burn-in")
    savefig(pb, joinpath(FIGDIR, "hw4_q1_burnin.png"))

    n_dist = stationary_distribution!(para, sols)
    mass = vec(sum(sols.μ, dims = 2))
    @assert all(sum(sols.T_star, dims = 2) .≈ 1.0) && sum(sols.μ) ≈ 1.0
    # Check both mass conservation and the lottery's mean placement.
    for j in 1:para.N_y, i in 1:para.N_a
        lo, hi, ω = lottery(sols.a_pol[i, j], para.a_grid)
        @assert isapprox(ω*para.a_grid[lo] + (1-ω)*para.a_grid[hi],
                         sols.a_pol[i, j]; atol = 1e-10)
    end
    println("     histogram: ", n_dist, " iterations;  at r = ", r0, ":")
    w_mc = fill(1/length(a_mc), length(a_mc))
    row("", "histogram", "panel")
    row("mean assets", rd(sum(mass .* para.a_grid), 5), rd(mean(a_mc), 5))
    row("top-10% share", rd(top_wealth_share(para.a_grid, mass, 0.1), 5),
                         rd(top_wealth_share(a_mc, w_mc, 0.1), 5))
    row("bottom-50% share", rd(1 - top_wealth_share(para.a_grid, mass, 0.5), 5),
                            rd(1 - top_wealth_share(a_mc, w_mc, 0.5), 5))
    row("mass at a_min", sg(mass[1]), sg(mean(a_mc .<= para.a_min + 1e-10)))
    row("mass at a_max", sg(mass[end]), sg(mean(a_mc .>= para.a_max - 1e-10)))

    pcdf = plot(para.a_grid, cumsum(mass); lw = 3, label = "histogram",
                xlabel = "a", ylabel = "P(assets ≤ a)", xlims = (0, 25),
                title = "Two roads to the same distribution")
    plot!(pcdf, para.a_grid, [mean(a_mc .<= x) for x in para.a_grid];
          lw = 3, ls = :dash, label = "panel, N = 10,000")
    savefig(pcdf, joinpath(FIGDIR, "hw4_q1_cdfs.png"))
    # Saving zero risks zero consumption next period, so Inada preferences
    # keep savings positive. The histogram's 2.6e-9 at zero is lottery error;
    # the simulated panel has no mass there.

    rs = collect(range(-0.01, 0.040; length = 15))
    A_hist  = [aggregate_savings(r, para, sols; method = :histogram) for r in rs]
    A_panel = [aggregate_savings(r, para, sols; method = :panel, y_panel = y_panel)
               for r in rs]
    println("     max |A_hist - A_panel| over the 15 rates = ",
            rd(maximum(abs.(A_hist .- A_panel))), " (relative ",
            rd(100*maximum(abs.(A_hist .- A_panel) ./ A_hist), 3), "%)")
    pA = plot(rs, A_hist; lw = 3, label = "A(r), histogram", xlabel = "r",
              ylabel = "capital", title = "Supply and demand for capital")
    plot!(pA, rs, A_panel; lw = 2, ls = :dash, label = "A(r), panel")
    plot!(pA, rs, [capital_demand(r, para) for r in rs]; lw = 3, label = "K^d(r)")
    vline!(pA, [1/para.β - 1]; ls = :dot, c = :black, label = "1/β - 1")
    savefig(pA, joinpath(FIGDIR, "hw4_q1_Ar.png"))
    # A(r) slopes UP and turns nearly vertical as r → 1/β - 1, where households
    # stop discounting relative to the market. K^d(r) slopes down; one crossing.

    # --- Part C -----------------------------------------------------------
    t0 = time(); r_star, K_star, para, sols = clear_market!(para, sols)
    t_root = time() - t0
    K_cm = K_complete_markets(para)
    mass = vec(sum(sols.μ, dims = 2))
    println("Q1C: root finder: r* = ", rd(r_star, 6), "  K* = ", rd(K_star, 5),
            "  (", rd(t_root, 1), " s)")
    println("     1/β - 1 = ", rd(1/para.β - 1, 6), " > r* ✓   K_complete_markets = ",
            rd(K_cm, 5), " < K* ✓   precautionary capital K*/K_cm - 1 = ",
            rd(100*(K_star/K_cm - 1), 2), "%")
    println("     K/Y = ", rd(K_star^(1 - para.α)), "   top-10% = ",
            rd(top_wealth_share(para.a_grid, mass, 0.1)), "   mass at a_max = ",
            sg(mass[end]))

    # The slope that decides which λ is safe, measured not guessed.
    hr = 1e-4
    A_slope = (aggregate_savings(r_star + hr, para, sols) -
               aggregate_savings(r_star - hr, para, sols)) / (2hr)
    s_slope = A_slope * para.α*(para.α - 1)*K_star^(para.α - 2)
    println("     dA/dr = ", rd(A_slope, 0), "  =>  d[A(r(K))]/dK = ", rd(s_slope, 2),
            ", so the damped map is stable only for λ < ", rd(2/(1 - s_slope), 3))
    _, _, _, _, n05, ok05, path05 = clear_market_damped!(para, sols; λ = 0.5,
                                                         max_iter = 24)
    println("     λ = 0.50: converged = ", ok05, ";  last K_n = ",
            rd.(path05[end-5:end], 3))
    paths = Dict{Float64, Vector{Float64}}(0.5 => path05)
    for λ in (0.2, 0.1)
        t0 = time()
        rλ, Kλ, _, _, nλ, okλ, pth = clear_market_damped!(para, sols; λ = λ,
                                                          max_iter = 200)
        paths[λ] = pth
        println("     λ = ", λ, ": ", okλ ? "converged" : "DID NOT converge", " in ",
                nλ, " steps, r = ", rd(rλ, 6), ", K = ", rd(Kλ, 5), "  (",
                rd(time() - t0, 1), " s)")
    end
    # λ = 0.5 cycles. Local stability requires |1-λ + λA'(r)r'(K)| < 1,
    # or λ < 0.156 here; λ = 0.1 converges.
    pdmp = plot(0:(length(path05)-1), path05; lw = 2, marker = :circle, ms = 2,
                label = "λ = 0.5 (cycles)", xlabel = "iteration n", ylabel = "K_n",
                title = "Damped capital update")
    plot!(pdmp, 0:(length(paths[0.1])-1), paths[0.1]; lw = 2, marker = :circle,
          ms = 2, label = "λ = 0.1 (converges)")
    hline!(pdmp, [K_star]; ls = :dash, c = :black, label = "K* from the root finder")
    savefig(pdmp, joinpath(FIGDIR, "hw4_q1_damped.png"))

    # --- Part D -----------------------------------------------------------
    # Lower job-finding probability, with mean income held fixed.
    para_d, sols_d = initialize(; p_ue = 0.50)
    r_d, K_d, para_d, sols_d = clear_market!(para_d, sols_d)
    mass_d = vec(sum(sols_d.μ, dims = 2))
    println("Q1D: p_ue = 0.50 -> u = ", rd(para_d.u_rate), ", y_e = ", rd(para_d.y_e),
            ", E[y] = ", rd(para_d.y_grid' * [1 - para_d.u_rate, para_d.u_rate]))
    println("     r* = ", rd(r_d, 6), " (was ", rd(r_star, 6), ")   K* = ", rd(K_d, 5),
            " (was ", rd(K_star, 5), ", ", rd(100*(K_d/K_star - 1), 2), "%)")
    println("     K/Y = ", rd(K_d^(1 - para_d.α)), " (was ", rd(K_star^(1 - para.α)),
            ")   top-10% = ", rd(top_wealth_share(para_d.a_grid, mass_d, 0.1)),
            " (was ", rd(top_wealth_share(para.a_grid, mass, 0.1)), ")")
    # More uninsurable risk raises precautionary saving: K rises and r falls.
    # A representative agent faces no idiosyncratic risk; its steady-state
    # Euler equation instead pins r = 1/β - 1.
    ph = plot(para.a_grid, mass; lw = 2, label = "p_ue = 0.75", xlims = (0, 25),
              xlabel = "a", ylabel = "mass",
              title = "Wealth distribution at equilibrium")
    plot!(ph, para_d.a_grid, mass_d; lw = 2, ls = :dash, label = "p_ue = 0.50")
    savefig(ph, joinpath(FIGDIR, "hw4_q1_distribution.png"))
    return r_star, K_star, para, sols
end

# =====================================================================
# Question 2: A multinomial logit
# =====================================================================
# u_ij = β x_j - α p_ij + ε_ij, u_i0 = ε_i0, ε Type-I extreme value, so
#   P_ij = exp(δ_ij) / (1 + Σ_k exp(δ_ik)),  δ_ij = β x_j - α p_ij.
# The data were simulated at (β, α) = (1.5, 1.0).

# choice_probs is pinned as (θ, p), so x comes in as a default keyword.
X_HW4 = [0.5, 1.0, 1.5, 2.0]

function load_choices(path)
    df = sort(CSV.read(path, DataFrame), [:consumer, :product])
    N, J = length(unique(df.consumer)), length(unique(df.product))
    # Sorted by consumer then product, the file is a J × N column-major block.
    x = [first(df.x[df.product .== j]) for j in 1:J]
    prices = permutedims(reshape(Vector{Float64}(df.price), J, N))
    picked = permutedims(reshape(Vector{Int}(df.chosen), J, N))
    chosen = [(k = findfirst(==(1), view(picked, i, :)); k === nothing ? 0 : k)
              for i in 1:N]
    return x, prices, chosen
end

function choice_probs(θ, p; x = X_HW4)
    β, α = θ[1], θ[2]
    δ = β .* x .- α .* p
    m = max(zero(eltype(δ)), maximum(δ))   # the outside option's δ is 0
    e = exp.(δ .- m)
    return vcat(exp(-m), e) ./ (exp(-m) + sum(e))
end

# Stable log-sum-exp likelihood, including the outside option's δ = 0.
# Preserve dual-number types for automatic differentiation.
function neg_loglik(θ, data)
    x, prices, chosen = data[1], data[2], data[3]
    β, α = θ[1], θ[2]
    T = typeof(β * α)
    N, J = size(prices)

    ll = zero(T)
    δ = Vector{T}(undef, J)
    for i in 1:N
        m = zero(T)
        for j in 1:J
            δ[j] = β*x[j] - α*prices[i, j]
            m = max(m, δ[j])
        end
        s = exp(-m)                        # the outside option
        for j in 1:J
            s += exp(δ[j] - m)
        end
        ll += (chosen[i] == 0 ? zero(T) : δ[chosen[i]]) - (m + log(s))
    end
    return -ll
end

# Estimate a = log α to enforce α > 0.
function estimate_mlogit(data; start = [0.0, 0.0])
    res = optimize(z -> neg_loglik([z[1], exp(z[2])], data), start, LBFGS();
                   autodiff = :forward)
    z = Optim.minimizer(res)
    return [z[1], exp(z[2])]
end

# Invert observed information in (β, log α) coordinates, then use
# the delta method: se(α̂) = α̂ se(log α̂).
function mlogit_se(data, θ̂)
    H = ForwardDiff.hessian(z -> neg_loglik([z[1], exp(z[2])], data),
                            [θ̂[1], log(θ̂[2])])
    V = inv(H)
    se_z = sqrt.(diag(V))
    return [se_z[1], θ̂[2]*se_z[2]], H, isposdef(H), V
end

mean_probs(θ, prices) =
    vec(mean(reduce(hcat, [choice_probs(θ, view(prices, i, :)) for i in axes(prices, 1)]);
             dims = 2))

function q2()
    path = joinpath(DATA_DIR, "choices.csv")

    # --- Part A -----------------------------------------------------------
    df = CSV.read(path, DataFrame)
    N = length(unique(df.consumer))
    println("Q2A: ", nrow(df), " rows, N = ", N, " consumers, J = ",
            length(unique(df.product)), ";  mean price = ", rd(mean(df.price)))
    by_prod = combine(groupby(df, :product), :price => mean => :mean_price,
                      :chosen => (c -> sum(c)/N) => :share)
    for r in eachrow(by_prod)
        println("     product ", r.product, ": x = ",
                first(df.x[df.product .== r.product]), "  mean price = ",
                rd(r.mean_price), "  share = ", rd(r.share))
    end
    println("     outside option share = ", rd(1 - sum(by_prod.share)))

    data = load_choices(path)
    x, prices, chosen = data
    @assert size(prices) == (N, 4) && length(chosen) == N && all(0 .<= chosen .<= 4)

    # --- Part B -----------------------------------------------------------
    θ_true = [1.5, 1.0]
    @assert sum(choice_probs(θ_true, view(prices, 1, :))) ≈ 1.0
    t0 = time(); θ̂ = estimate_mlogit(data); t_est = time() - t0
    println("Q2B: -log L at the truth = ", rd(neg_loglik(θ_true, data)),
            ";  θ̂ = (β̂ = ", rd(θ̂[1]), ", α̂ = ", rd(θ̂[2]), "), -log L = ",
            rd(neg_loglik(θ̂, data)), "  (", rd(t_est, 2), " s)")
    # Estimates (1.563, 1.034) recover (1.5, 1.0) within one standard error.

    # --- Part C -----------------------------------------------------------
    se, H, pd, V = mlogit_se(data, θ̂)
    println("Q2C: Hessian positive definite: ", pd, "  (eigenvalues ",
            rd.(eigvals(H), 1), ");  corr(β̂, α̂) = ", rd(V[1,2]/sqrt(V[1,1]*V[2,2]), 3))
    for (k, name, truth) in ((1, "β̂", 1.5), (2, "α̂", 1.0))
        println("     ", name, " = ", rd(θ̂[k]), " (se ", rd(se[k]), ")   truth ",
                truth, " is ", rd(abs(θ̂[k] - truth)/se[k], 2), " se away ",
                abs(θ̂[k] - truth) < 2se[k] ? "✓" : "✗")
    end
    # The information matrix is positive definite. Estimates correlate at
    # +0.93 because price and x are correlated in the DGP (p = 2 + 0.5x + ε).

    # --- Part D -----------------------------------------------------------
    base = mean_probs(θ̂, prices)
    emp = [mean(chosen .== j) for j in 0:4]
    prices_up = copy(prices); prices_up[:, 4] .*= 1.10
    new = mean_probs(θ̂, prices_up)
    println("Q2D: option     observed  predicted   after +10% on product 4")
    for j in 0:4
        println("     ", j == 0 ? "outside " : "prod. $j ", lpad(rd(emp[j+1]), 9),
                lpad(rd(base[j+1]), 11), lpad(rd(new[j+1]), 11), "  (",
                rd(100*(new[j+1]/base[j+1] - 1), 3), "%)")
    end
    println("     max |predicted - observed| = ", rd(maximum(abs.(base .- emp))),
            ";  own-price elasticity of product 4 ≈ ",
            rd((new[5]/base[5] - 1)/0.10, 3))
    # Two parameters do not generally fit every observed share exactly.
    # Product 4 loses 19%; other shares rise roughly 9%. IIA implies proportional
    # substitution within each consumer, not necessarily after averaging over
    # their different prices. It cannot favor close substitutes beyond their
    # existing probabilities (e.g. another car versus a bicycle). Nested or
    # mixed logit can relax IIA; adding characteristics alone does not.
    return θ̂, se
end

# =====================================================================
# Question 3: SMM on the Question 1 model
# =====================================================================
# Fit (β, γ) using Q1's model; add the bottom-50% share in part C.

# K/Y = ā^(1-α) (mean wealth ā, since Y = K^α with L = 1). top_wealth_share
# only sees values and weights, so it works unchanged on raw microdata.
function data_moments(path = joinpath(DATA_DIR, "hw4_wealth.csv");
                      α = AiyagariParameters().α)
    a = Float64.(CSV.read(path, DataFrame).wealth)
    m = fill(1/length(a), length(a))
    return [mean(a)^(1 - α), top_wealth_share(a, m, 0.1), 1 - top_wealth_share(a, m, 0.5)]
end

# Return all three moments. Reuse sols to warm-start nearby parameter values;
# shocks supplies fixed uniforms for the panel method only.
function model_moments(θ; shocks = nothing, sols = nothing, method = :histogram,
                       kwargs...)
    β, γ = θ[1], θ[2]
    para, s = initialize(; β = β, γ = γ, kwargs...)
    sols === nothing || (s = sols)                          # warm start
    r_star, K_star, para, s = clear_market!(para, s; method = method, shocks = shocks,
                                            r_guess = s.r, r_hi = 1/β - 1)
    a, m = wealth_distribution(para, s; method = method)
    return [K_star^(1 - para.α), top_wealth_share(a, m, 0.1),
            1 - top_wealth_share(a, m, 0.5)], para, s
end

# m_data carries 2 moments (parts A/B/D) or 3 (part C), so slice to its length
# rather than branching on which part is calling.
function smm_objective(θ, m_data, W; shocks = nothing, kwargs...)
    # A guard, not a constraint: Nelder-Mead is unconstrained and will try
    # β ≥ 1, where the household problem has no stationary solution.
    (0.5 < θ[1] < 0.9999 && θ[2] > 0) || return 1e6
    m, _, _ = model_moments(θ; shocks = shocks, kwargs...)
    g = m_data .- m[1:length(m_data)]
    return g' * W * g
end

# Transform unconstrained coordinates to β ∈ (0.85, 0.995) and γ > 0.
logistic(z) = 1 / (1 + exp(-z))
logit(p) = log(p / (1 - p))
β_LO, β_HI = 0.85, 0.995
θ_of_z(z) = [β_LO + (β_HI - β_LO)*logistic(z[1]), exp(z[2])]
z_of_θ(θ) = [logit((θ[1] - β_LO)/(β_HI - β_LO)), log(θ[2])]

function estimate_smm(m_data, W; shocks = nothing, θ0 = [0.96, 2.0],
                      method = :histogram, N_a = 100, iterations = 500, kwargs...)
    # One sols object for the whole estimation, so every solve warm-starts from
    # the last. Its grid never changes, only β and γ do.
    _, s = initialize(; N_a = N_a, kwargs...)
    n_solves = 0
    f(z) = (n_solves += 1;
            smm_objective(θ_of_z(z), m_data, W; shocks = shocks, sols = s,
                          method = method, N_a = N_a, kwargs...))
    t0 = time()
    res = optimize(f, z_of_θ(θ0), NelderMead(),
                   Optim.Options(x_abstol = 1e-5, f_abstol = 1e-12,
                                 iterations = iterations))
    return θ_of_z(Optim.minimizer(res)), n_solves, time() - t0, res, s
end

function q3()
    # --- Part A: data_moments / model_moments / smm_objective, defined above --
    m_data = data_moments()
    println("Q3A: data moments: K/Y = ", rd(m_data[1]), ", top-10% = ",
            rd(m_data[2]), ", bottom-50% = ", rd(m_data[3]))

    # --- Part B -----------------------------------------------------------
    # Cheap solves: N_a = 100 rather than 300, the histogram rather than the
    # panel, and one warm-started sols across the whole run.
    N_a = 100
    W2 = Matrix(1.0I, 2, 2)
    θ̂, n_solves, t_smm, res, _ = estimate_smm(m_data[1:2], W2; N_a = N_a)
    m̂, _, _ = model_moments(θ̂; N_a = N_a)
    println("Q3B: θ̂ = (β̂ = ", rd(θ̂[1], 5), ", γ̂ = ", rd(θ̂[2], 5), ") from ",
            n_solves, " model solves in ", rd(t_smm, 1), " s (",
            rd(t_smm/n_solves, 2), " s each), objective ", sg(Optim.minimum(res)))
    println("     fitted: K/Y = ", rd(m̂[1]), " (data ", rd(m_data[1]),
            "), top-10% = ", rd(m̂[2]), " (data ", rd(m_data[2]), ")")
    # This run fits both moments nearly exactly, but β̂ is 0.27% and γ̂ 7.2%
    # from the hidden truth (0.94, 3.0). Two parameters and two moments allow
    # an exact fit only when the target is attainable and the optimizer finds it.

    # θ̂ absorbs whatever bias the coarse grid leaves: re-solve at θ̂ on a finer
    # grid and the moments move.
    m_fine, _, _ = model_moments(θ̂; N_a = 300)
    println("     same θ̂ at N_a = 300: K/Y = ", rd(m_fine[1]), ", top-10% = ",
            rd(m_fine[2]), ";  |moment shift| = (", rd(abs(m_fine[1] - m̂[1])),
            ", ", rd(abs(m_fine[2] - m̂[2])), ")")

    # Histogram moments have grid kinks. Use central differences with steps
    # wide enough to avoid the misleading sign from a small one-sided step.
    h = [0.005, 0.05]
    J = zeros(2, 2)
    for k in 1:2
        θp = copy(θ̂); θp[k] += h[k]
        θm = copy(θ̂); θm[k] -= h[k]
        mp, _, _ = model_moments(θp; N_a = N_a)
        mm, _, _ = model_moments(θm; N_a = N_a)
        J[:, k] = (mp[1:2] .- mm[1:2]) ./ (2h[k])
    end
    E = [J[i, k] * θ̂[k] / m̂[i] for i in 1:2, k in 1:2]     # elasticities
    println("     ∂m/∂θ'        w.r.t. β      w.r.t. γ    [same as elasticities]")
    println("       K/Y     ", lpad(rd(J[1,1], 3), 12), lpad(rd(J[1,2], 4), 14),
            "  [", rd(E[1,1], 2), ", ", rd(E[1,2], 3), "]")
    println("       top-10% ", lpad(rd(J[2,1], 5), 12), lpad(rd(J[2,2], 5), 14),
            "  [", rd(E[2,1], 2), ", ", rd(E[2,2], 3), "]")
    println("     |J| = ", sg(det(J)), "   cond(J) = ", rd(cond(J), 1),
            "   cond(E) = ", rd(cond(E), 1))
    # Elasticities remove units: cond(E) ≈ 20 versus cond(J) ≈ 1228.
    # Both moments respond more to β; γ is identified by the remaining change
    # in concentration. This weak sensitivity amplifies coarse-grid error:
    # γ̂ falls from 3.217 at N_a = 100 to 3.031 at 300 and 2.984 at 500.

    # --- Part C -----------------------------------------------------------
    θ̂3, n3, t3, res3, _ = estimate_smm(m_data, Matrix(1.0I, 3, 3); N_a = N_a)
    m̂3, _, _ = model_moments(θ̂3; N_a = N_a)
    obj3 = Optim.minimum(res3)
    println("Q3C: θ̂ = (β̂ = ", rd(θ̂3[1], 5), ", γ̂ = ", rd(θ̂3[2], 5), ") from ", n3,
            " solves in ", rd(t3, 1), " s, objective = ", sg(obj3, 4))
    println("     fitted: K/Y = ", rd(m̂3[1]), ", top-10% = ", rd(m̂3[2]),
            ", bottom-50% = ", rd(m̂3[3]), "   residuals ", rd.(m_data .- m̂3))
    @assert !isapprox(obj3, 0.0; atol = 1e-8)   # confirm it is NOT exactly zero
    # Objective ≈ 5.8e-7, with top-10% fitting worst. K/Y largely pins β,
    # leaving γ to fit two wealth shares that are not jointly attainable here.

    # --- Part D -----------------------------------------------------------
    m_us = [3.0, 0.76]
    # Limit the search to 60 iterations; report the best fit found.
    θ_us, n_us, t_us, res_us, _ = estimate_smm(m_us, W2; N_a = N_a, θ0 = θ̂3[1:2],
                                               iterations = 60)
    m_fit, _, _ = model_moments(θ_us; N_a = N_a)
    println("Q3D: targets K/Y = ", m_us[1], ", top-10% = ", m_us[2], " -> θ̂ = (β̂ = ",
            rd(θ_us[1], 5), ", γ̂ = ", rd(θ_us[2], 5), ") after ", n_us, " solves, ",
            rd(t_us, 1), " s")
    println("     best fit: K/Y = ", rd(m_fit[1]), ", top-10% = ", rd(m_fit[2]),
            ", objective ", rd(Optim.minimum(res_us)), " (vs ",
            sg(Optim.minimum(res)), " in part B)")

    # Check a coarse parameter grid for comparison with the optimizer's result.
    _, s_sweep = initialize(; N_a = 40)
    best = (-Inf, 0.0, 0.0)
    for β in (0.90, 0.95, 0.99), γ in (0.1, 0.5, 1.0, 2.0, 4.0)
        m, _, s_sweep = model_moments([β, γ]; sols = s_sweep, N_a = 40)
        m[2] > best[1] && (best = (m[2], β, γ))
    end
    println("     highest top-10% share on a coarse (β, γ) sweep: ", rd(best[1]),
            " at (β = ", best[2], ", γ = ", best[3], ") -- the US target ",
            m_us[2], " is a factor of ", rd(m_us[2]/best[1], 1), " away")
    # β̂ ≈ 0.962 matches K/Y ≈ 3, while γ̂ falls to ≈ 0.020 trying to raise
    # concentration. Top-10% reaches only 0.200 versus 0.76; the coarse sweep
    # reaches 0.205, which is evidence of poor fit, not a global upper bound.
    # The model omits persistent income differences, entrepreneurial and return
    # risk, bequests, and heterogeneous patience that can concentrate wealth.
    return θ̂, m̂
end

# ---------------------------------------------------------------------
if abspath(PROGRAM_FILE) == @__FILE__
    q1()
    q2()
    q3()
end
