# homework3_solution.jl
# Computational Bootcamp, Summer 2026 -- Homework 3 solutions
#
# Q1  Parallelization and the speed of convergence
# Q2  Job search with an expiring benefit (finite-horizon DP)
# Q3  Endowment economy: Howard iteration and the endogenous grid method
#
#     julia --project=. homework3.jl

using Distributed, Plots, Optim

if nprocs() == 1
    addprocs(min(8, Sys.CPU_THREADS))
end

# @everywhere covers the master too, so this is Q2/Q3's import as well.
@everywhere using Interpolations, Roots, Optim

FIGDIR = joinpath(@__DIR__, "figures")

# =====================================================================
# Question 1: Parallelization and the speed of convergence
# =====================================================================

# --- Part A: the model, the steady state, and the half-life ----------------
# Anything the workers call must be defined on the workers, so all of part A
# goes in one @everywhere block (which also defines it on the master).
@everywhere begin
    α_Q1 = 0.36
    δ_Q1 = 0.025

    # Closed form from the Euler equation at k' = k: 1 = β(α k^(α-1) + 1 - δ).
    # Used only to CHECK steady_state(), never to compute it.
    steady_state_closed(β; α = α_Q1, δ = δ_Q1) = ((1/β - 1 + δ)/α)^(1/(α - 1))

    # VFI on Homework 1's deterministic model, but with a CONTINUOUS choice of
    # k': cubic-spline V, Brent per state. A grid search would quantize the
    # policy into a dead band around k_ss and wreck the root-find below.
    # kmax = 75 because k_ss(0.995) ≈ 48.5.
    function solve_growth(β; α = α_Q1, δ = δ_Q1, nk = 200, kmin = 0.5, kmax = 75.0,
                          tol = 1e-6, maxiter = 20_000)
        k_grid = range(kmin, kmax; length = nk)
        V = zeros(nk); V_next = zeros(nk); pol = zeros(nk)
        diff = Inf; n = 0
        while diff > tol && n < maxiter
            n += 1
            Vf = cubic_spline_interpolation(k_grid, V; extrapolation_bc = Line())
            @inbounds for i in 1:nk
                budget = k_grid[i]^α + (1 - δ)*k_grid[i]
                hi = min(budget - 1e-8, kmax)             # can't consume exactly 0
                obj(kp) = -(log(budget - kp) + β*Vf(kp))
                res = optimize(obj, kmin, hi, Brent())
                V_next[i] = -Optim.minimum(res); pol[i] = Optim.minimizer(res)
            end
            diff = maximum(abs.(V_next .- V))
            V .= V_next
        end
        return collect(k_grid), pol
    end

    # k_ss from the model itself: root-find g(k) - k = 0 on the interpolated policy.
    function policy_fixed_point(k_grid, pol)
        g = cubic_spline_interpolation(range(k_grid[1], k_grid[end]; length = length(k_grid)),
                                        pol; extrapolation_bc = Line())
        find_zero(k -> g(k) - k, (k_grid[1], k_grid[end])), g
    end

    function steady_state(β; kwargs...)
        k_grid, pol = solve_growth(β; kwargs...)
        kss, _ = policy_fixed_point(k_grid, pol)
        return kss
    end

    # Periods to close half the gap to k_ss from k_ss/2. Solves the model once
    # and reuses the policy for both k_ss and the iteration -- this is the
    # function swept over 60 β's below.
    function half_life(β; kwargs...)
        k_grid, pol = solve_growth(β; kwargs...)
        kss, g = policy_fixed_point(k_grid, pol)
        k = kss/2
        gap0 = abs(kss - k)
        t = 0
        while abs(kss - k) > gap0/2 && t < 10_000
            t += 1
            k = g(k)
        end
        return t
    end
end

# --- Part B: the sweep, serial and distributed -----------------------------
sweep_serial(βs) = [half_life(β) for β in βs]

# pmap, not @distributed: the tasks are very unequal (β = 0.995 costs ~23x
# β = 0.90) so dynamic scheduling matters, and each is far too big for the
# per-task communication to. The WorkerPool gives the 1/2/4/all-worker table
# without adding and removing processes.
sweep_pmap(βs, nw = nworkers()) = pmap(half_life, WorkerPool(workers()[1:nw]), βs)

function q1()
    βs = collect(range(0.90, 0.995; length = 60))

    # Part A check: k_ss from the policy vs the closed form, over the whole sweep.
    err = maximum(abs(steady_state(β) - steady_state_closed(β)) for β in βs)
    println("Q1A: max |k_ss (policy fixed point) - k_ss (closed form)| over the β grid = ",
            round(err, sigdigits = 3))
    for β in (0.90, 0.995)
        println("     β = ", β,
                "  k_ss (policy) = ", round(steady_state(β), digits = 4),
                "   k_ss (closed form) = ", round(steady_state_closed(β), digits = 4))
    end

    half_life(0.95); sweep_pmap(βs[1:nworkers()])      # compile everywhere

    t0 = time(); hl = sweep_serial(βs); t_ser = time() - t0
    println("Q1B: serial ", round(t_ser, digits = 2), " s")
    for nw in unique((1, 2, 4, nworkers()))
        nw > nworkers() && continue
        t0 = time(); hl_p = sweep_pmap(βs, nw); t = time() - t0
        @assert hl_p == hl                             # identical, not just close
        println("     pmap, ", nw, " worker(s): ", round(t, digits = 2),
                " s   speedup ", round(t_ser/t, digits = 2), "x")
    end

    p = plot(βs, hl; lw = 2, legend = false,
             xlabel = "β", ylabel = "periods to close half the gap",
             title = "Convergence half-life")
    savefig(p, joinpath(FIGDIR, "hw3_q1_halflife.png"))

    println("Q1C: half-life rises from ", hl[1], " periods at β = ", βs[1],
            " to ", hl[end], " at β = ", βs[end])
    # Interpretation. A MORE patient economy converges more SLOWLY: 7 periods at
    # β = 0.90 vs 24 at β = 0.995. Patience raises k_ss but also makes the
    # household content to spread the transition out (the stable eigenvalue of
    # the linearized policy → 1 as β → 1). This is the neoclassical model's
    # well-known slow-convergence problem.
    #
    # Timings (8 performance cores, M1 Pro; ±several percent run to run):
    #   1 worker ~4.8 s (1.0x) | 2 ~2.6 s (1.9x) | 4 ~1.5 s (3.2x) | 8 ~1.0 s (4.6x)
    #
    # Why not 8x? Ideal is serial/8 = 0.60 s, actual ~1.02 s. Timing each task
    # and tagging it with its worker splits the 0.42 s gap roughly two-to-one:
    #   - Load imbalance (~0.28 s): the 23x spread in task cost, dealt in the
    #     order of βs, leaves the busiest worker ~1.01 s against a ~0.73 s
    #     average. Sorting βs longest-first recovers most of this.
    #   - Core contention (~0.13 s): worker busy time totals ~5.8 s vs ~4.8 s
    #     serial -- 9 processes on 8 cores and one memory bus.
    # NOT communication (one Float64 out, one Int back; the 1-worker pmap costs
    # only ~3% over the serial loop), and not the longest task (~0.53 s, still
    # under serial/8 -- that bound would only bind past ~9 workers).
    return hl
end

# =====================================================================
# Question 2: Job search with an expiring benefit
# =====================================================================
# Offers are iid on a 101-point grid; accepting w is worth w/(1-β) forever.
WGRID = collect(range(10.0, 60.0; length = 101))
PROB  = fill(1/101, 101)

# Backward induction on the value of REJECTING. R_t doesn't depend on today's
# offer, so the policy is one reservation wage per period, w̄_t = (1-β) R_t,
# and there is no optimizer anywhere in this problem.
function reservation_path(T; β = 0.98, b = 25.0, w = WGRID, p = PROB)
    R = zeros(T)
    # Terminal condition: in T+1 the benefit is gone and any offer is accepted,
    # so E[V_{T+1}] = E[w]/(1-β) -- known, hence no fixed point.
    R[T] = b + β * sum(p .* w)/(1 - β)
    for t in (T-1):-1:1
        EV = sum(p .* max.(w ./ (1 - β), R[t+1]))
        R[t] = b + β * EV
    end
    return (1 - β) .* R
end

# Infinite horizon: same equation, but R is its own continuation -- a scalar
# fixed point, solved with a root finder.
function reservation_infinite(; β = 0.98, b = 25.0, w = WGRID, p = PROB)
    f(R) = b + β * sum(p .* max.(w ./ (1 - β), R)) - R
    R = find_zero(f, (b/(1 - β), maximum(w)/(1 - β)))
    return (1 - β) * R
end

function q2()
    T = 40
    wbar = reservation_path(T)
    println("Q2A/B: reservation wage as the benefit runs out")
    for t in (1, 10, 20, 30, 38, 39, 40)
        println("       t = ", t, " (", T - t + 1, " periods of benefit left): w̄ = ",
                round(wbar[t], digits = 4))
    end
    p = plot(1:T, wbar; lw = 2, legend = false, xlabel = "period t",
             ylabel = "reservation wage", title = "Reservation wage, T = 40")
    savefig(p, joinpath(FIGDIR, "hw3_q2_reservation.png"))
    # Nearly flat for most of the spell, then falls off a cliff in the last few
    # periods: early on there are many draws left, so the option value of
    # rejecting is close to its infinite-horizon value; as exhaustion nears that
    # option is worth less and the worker accepts wages they'd have refused a
    # month earlier. Hence the job-finding spike at benefit exhaustion.

    wbar_inf = reservation_infinite()
    println("Q2C: infinite-horizon w̄ = ", round(wbar_inf, digits = 6))
    for TT in (10, 40, 100, 400)
        w1 = reservation_path(TT)[1]
        println("     T = ", TT, ": w̄_1 = ", round(w1, digits = 6),
                "   gap = ", round(abs(w1 - wbar_inf), sigdigits = 3))
    end
    # The gap collapses (1.6 at T=10, 6e-3 at T=40, 1e-7 at T=100, machine
    # precision by T=400). Backward induction is the Bellman operator applied T
    # times to a fixed terminal condition; the operator is a contraction, so it
    # converges geometrically to the same fixed point regardless of where it
    # starts -- long-horizon backward induction and VFI are the same computation.

    println("     benefit generosity:")
    for b in (10.0, 20.0, 30.0, 40.0)
        wb = reservation_infinite(; b = b)
        pacc = sum(PROB .* (WGRID .>= wb))
        println("       b = ", b, "  w̄ = ", round(wb, digits = 3),
                "  P(accept) = ", round(pacc, digits = 4),
                "  duration = ", round(1/pacc, digits = 2), " periods")
    end
    # A more generous benefit raises w̄, so the worker is choosier and stays
    # unemployed longer (5.6 periods at b=10, 9.2 at b=40). The trade-off:
    # benefits insure consumption and buy time to find a better match, but
    # lengthen unemployment -- the core tension in the empirical UI literature.
    return wbar
end

# =====================================================================
# Question 3: Endowment economy -- Howard iteration and EGM
# =====================================================================
Base.@kwdef struct Endowment
    γ::Float64 = 2.0                                  # CRRA
    r::Float64 = 0.03
    β::Float64 = 0.96
    Y::Vector{Float64} = [1.0, 0.5]                   # y_h, y_l
    Π::Matrix{Float64} = [0.95 0.05; 0.75 0.25]       # Π[i,j] = P(y'=Y[j] | y=Y[i])
    na::Int = 101
    # @kwdef defaults may refer to earlier fields, so the grid always matches na.
    a_grid::Vector{Float64} = collect(range(0.0, 5.0; length = na))  # a_grid[1] = 0 is the borrowing limit
    tol::Float64 = 1e-5
    maxiter::Int = 5_000
end

util(c, γ) = γ == 1 ? log(c) : c^(1 - γ)/(1 - γ)
u_prime(c, γ) = c^(-γ)                                # u'(c)
u_prime_inv(x, γ) = x^(-1/γ)                          # (u')^{-1}, closed form for CRRA

# VFI with a continuous choice of a' (Brent on the interpolated V). `m` is the
# number of Howard policy-evaluation sweeps after each maximizing sweep; m = 0
# is plain VFI. kwargs... forwards into Endowment, so any field can be
# overridden by name.
function solve_vfi(; m = 0, kwargs...)
    p = Endowment(; kwargs...)
    (; γ, r, β, Y, Π, na, a_grid, tol, maxiter) = p
    ny = length(Y)
    amin, amax = a_grid[1], a_grid[end]                # amin = borrowing limit
    V = zeros(na, ny); V_next = zeros(na, ny)
    pol = zeros(na, ny)                               # policy a'(a,y)
    ug  = zeros(na, ny)                               # flow payoff at that policy
    diff = Inf; n = 0
    while diff > tol && n < maxiter
        n += 1
        Vf = [linear_interpolation(a_grid, V[:, yi]; extrapolation_bc = Line()) for yi in 1:ny]
        @inbounds for yi in 1:ny
            EV(ap) = sum(Π[yi, yp] * Vf[yp](ap) for yp in 1:ny)
            for i in 1:na
                budget = (1 + r)*a_grid[i] + Y[yi]
                hi = min(budget - 1e-10, amax)         # can't consume exactly 0
                obj(ap) = -(util(budget - ap, γ) + β*EV(ap))
                res = optimize(obj, amin, hi, Brent())
                ap_star = Optim.minimizer(res)
                pol[i, yi] = ap_star
                ug[i, yi]  = util(budget - ap_star, γ)
                V_next[i, yi] = -Optim.minimum(res)
            end
        end
        diff = maximum(abs.(V_next .- V))
        V .= V_next

        # Howard: re-apply the policy just found m more times. No optimizer in
        # here -- that is the whole point, and why it is nearly free.
        for _ in 1:m
            Vf = [linear_interpolation(a_grid, V[:, yi]; extrapolation_bc = Line()) for yi in 1:ny]
            @inbounds for yi in 1:ny
                for i in 1:na
                    EV = sum(Π[yi, yp] * Vf[yp](pol[i, yi]) for yp in 1:ny)
                    V_next[i, yi] = ug[i, yi] + β*EV
                end
            end
            V .= V_next
        end
    end
    return V, pol, n
end

# EGM. The budget constraint is linear in a, so no change of variables is
# needed: given c and a', today's assets are a = (c + a' - y)/(1+r).
function solve_egm(; kwargs...)
    p = Endowment(; kwargs...)
    (; γ, r, β, Y, Π, na, a_grid, tol, maxiter) = p
    ny = length(Y)
    apj = a_grid                                      # grid over TOMORROW's assets a'

    c = [(1 + r)*a_grid[i] + Y[yi] for i in 1:na, yi in 1:ny]   # guess: consume everything
    c_new = similar(c)
    diff = Inf; n = 0
    while diff > tol && n < maxiter
        n += 1
        cf = [linear_interpolation(a_grid, c[:, yi]; extrapolation_bc = Line()) for yi in 1:ny]

        for yi in 1:ny
            # Step 1: Euler RHS at each a'_j, then invert u' in closed form for
            # consumption endogenous to that a'_j.
            RHS = zeros(na)
            for j in 1:na, yp in 1:ny
                RHS[j] += Π[yi, yp] * u_prime(cf[yp](apj[j]), γ)
            end
            RHS .*= β*(1 + r)
            c_endog = u_prime_inv.(RHS, γ)
            # Step 2: today's assets supporting (c_endog[j], a'_j) -- the
            # endogenous grid, which differs across y.
            a_endog = (c_endog .+ apj .- Y[yi]) ./ (1 + r)

            # Step 3: interpolate back onto the fixed grid.
            c_of_a = linear_interpolation(a_endog, c_endog; extrapolation_bc = Line())
            for i in 1:na
                c_new[i, yi] = c_of_a(a_grid[i])
            end
            # Step 4: where that implies a' < 0 the constraint binds -- set
            # a' = 0 and read c off the budget line. This puts in the kink.
            for i in 1:na
                ap_implied = (1 + r)*a_grid[i] + Y[yi] - c_new[i, yi]
                if ap_implied < 0
                    c_new[i, yi] = (1 + r)*a_grid[i] + Y[yi]
                end
            end
        end
        diff = maximum(abs.(c_new .- c))
        c .= c_new
    end

    # Savings policy from the budget constraint, folded in here so the caller
    # never has to rebuild the params struct to recover a'.
    ap = zeros(na, ny)
    for yi in 1:ny, i in 1:na
        ap[i, yi] = (1 + r)*a_grid[i] + Y[yi] - c[i, yi]
    end
    return a_grid, ap, c, n
end

function q3()
    p = Endowment()                # defaults, used below only to read grid values
    solve_vfi(; na = 21)  # compile

    t0 = time(); V, pol, n = solve_vfi(); t_vfi = time() - t0
    println("Q3A: plain VFI: ", n, " sweeps, ", round(t_vfi, digits = 2), " s")

    pv = plot(p.a_grid, V; label = ["y high" "y low"], lw = 2,
              xlabel = "a", ylabel = "V", title = "Value functions")
    savefig(pv, joinpath(FIGDIR, "hw3_q3_value.png"))

    pp = plot(p.a_grid, pol; label = ["y high" "y low"], lw = 2,
              xlabel = "a", ylabel = "a'", title = "Policy functions (VFI)")
    plot!(pp, p.a_grid, p.a_grid; label = "45°", ls = :dash, c = :black)
    savefig(pp, joinpath(FIGDIR, "hw3_q3_policy.png"))
    # Two things to notice. (1) The low-income policy is pinned flat at a' = 0
    # for a ≤ 0.20: the borrowing constraint binding. (2) Both policies lie
    # below the 45-degree line almost everywhere; the high-income one crosses
    # once, at a ≈ 0.77 -- the target wealth a household drifts toward while it
    # stays in that state (saving below, dissaving above). The low-income
    # target is the constraint itself. Behind this: β(1+r) = 0.9888 < 1, i.e.
    # r = 0.03 below the discount rate (1-β)/β ≈ 0.0417, so assets do not grow
    # without bound -- the Huggett/Aiyagari condition for a stationary wealth
    # distribution, with the income shock knocking households between targets.

    println("Q3B: Howard iteration")
    for m in (0, 5, 20, 50, 100)
        t0 = time(); _, _, nm = solve_vfi(; m = m); el = time() - t0
        println("       m = ", m, ": ", nm, " maximizing sweeps  ",
                round(el, digits = 2), " s  (", round(t_vfi/el, digits = 1), "x)")
    end
    # 285 sweeps / ~1.2 s at m=0 down to 11 sweeps / ~0.10 s at m=50: ~12x for
    # five extra lines. The gain flattens because the floor is the number of
    # sweeps the POLICY needs to settle (~11 here) -- Howard removes the wait
    # for the VALUE to converge (geometric at rate β) but every policy change
    # still costs a maximizing sweep. Past m ≈ 50 you re-evaluate an already
    # converged policy, and m=100 is slightly slower than m=50.

    solve_egm(; na = 21)  # compile
    t0 = time(); _, ap_egm, _, ne = solve_egm(); t_egm = time() - t0
    println("Q3C: EGM: ", ne, " iterations, ", round(t_egm, digits = 3), " s")

    d = abs.(ap_egm .- pol)
    near = (p.a_grid .>= 0.2) .& (p.a_grid .<= 1.0)    # just above the low-y kink
    println("     max|a'_EGM - a'_VFI| = ", round(maximum(d), digits = 4),
            "   (near kink: ", round(maximum(d[near, :]), digits = 4),
            ")   grid step = ", round(p.a_grid[2] - p.a_grid[1], digits = 4))
    # EGM has no optimizer and no value function to search over -- each sweep is
    # a closed-form inversion plus one interpolation -- so ~0.001 s against
    # plain VFI's ~1.1 s, three orders of magnitude, far more than Howard bought.

    for (yi, lbl) in enumerate(("y high", "y low"))
        j = findfirst(>(1e-10), ap_egm[:, yi])
        astar = j === nothing ? p.a_grid[end] : p.a_grid[j]
        println("     a*(", lbl, ") ~ ", round(astar, digits = 4),
                "  (EGM constraint threshold)")
    end
    # a*(y high) = 0: the high-income household saves even with no assets, so
    # the constraint never binds. The number printed for y low is the first
    # UNconstrained grid point, so true a*(y low) ∈ (0.20, 0.25]. VFI lands on
    # the same grid point in both cases -- both methods pin a' = 0 there by
    # construction, so they must agree on where the constraint binds.

    pe = plot(p.a_grid, pol; label = ["VFI, y high" "VFI, y low"], lw = 3,
              xlabel = "a", ylabel = "a'", title = "Policy: VFI vs EGM")
    plot!(pe, p.a_grid, ap_egm; label = ["EGM, y high" "EGM, y low"], ls = :dash, lw = 2)
    plot!(pe, p.a_grid, p.a_grid; label = "45°", ls = :dot, c = :black)
    savefig(pe, joinpath(FIGDIR, "hw3_q3_egm_policy.png"))
    # They disagree most just above a*(y low), where V has a kink that VFI must
    # search over while EGM only interpolates a smooth consumption function.
    # Trust EGM there. Away from the kink (a > 1) they agree about twice as
    # tightly (median gap 7e-4 vs 1.5e-3 just above a*).
    return V, pol, ap_egm
end

# ---------------------------------------------------------------------
if abspath(PROGRAM_FILE) == @__FILE__
    isdir(FIGDIR) || mkpath(FIGDIR)
    q1(); q2(); q3()
end
