#!/usr/bin/env julia using CSV using DataFrames using Statistics using Printf const ECON_VARS = Set(["lrecon_ches", "lrecon_poppa", "lrecon_gps", "lrecon_vparty", "welf_vparty"]) const CULT_VARS = Set(["galtan_ches", "libcon_gps", "immig_vparty", "lgbt_vparty", "culsup_vparty", "relig_vparty", "gender_vparty"]) function region_for(country::AbstractString) europe = Set(["AL","AT","BA","BE","BG","BY","CH","CY","CZ","DE","DK","EE","ES","FI","FR","GB","GE","GR","HR","HU","IE","IS","IT","LT","LU","LV","MD","ME","MK","MT","NL","NO","PL","PT","RO","RS","RU","SE","SI","SK","TR","UA"]) latin = Set(["AR","BO","BR","CL","CO","CR","DO","EC","MX","PA","PE","UY"]) country in europe && return "Europe" country in latin && return "Latin America" country in Set(["CA", "US"]) && return "North America" country in Set(["AU", "JP", "KR", "LK", "NZ"]) && return "Asia-Pacific" return "Other" end function wilson(covered::Int, n::Int; z=1.96) n == 0 && return (missing, missing) p = covered / n denom = 1 + z^2 / n centre = (p + z^2 / (2n)) / denom half = z * sqrt(p * (1-p) / n + z^2 / (4n^2)) / denom return (centre - half, centre + half) end function safe_cor(x, y) length(x) < 3 && return missing std(x) == 0 && return missing std(y) == 0 && return missing return cor(x, y) end function summarize_group(d::AbstractDataFrame, group_type::String, group::String) n = nrow(d) overlapping = count(d.within_latent_interval_95) lo, hi = wilson(overlapping, n) DataFrame( dimension = first(d.dimension), group_type = group_type, group = group, n = n, parties = length(unique(d.party_id)), pearson_r = safe_cor(d.expert_value, d.model_value), mae = mean(d.absolute_error), rmse = sqrt(mean(d.error .^ 2)), bias_expert_minus_model = mean(d.error), latent_interval_overlap_95 = overlapping / n, latent_interval_overlap_ci_lower = lo, latent_interval_overlap_ci_upper = hi, mean_interval_width = mean(d.interval_upper .- d.interval_lower), mean_nearest_text_distance = mean(d.nearest_text_distance) ) end function main() length(ARGS) >= 1 || error("Usage: summarize_blocked_validation.jl MODEL_OUTPUT [RUN_DIR] [OUTPUT_DIR]") model_file = ARGS[1] run_dir = length(ARGS) >= 2 ? ARGS[2] : "_local/validation/blocked_party" output_dir = length(ARGS) >= 3 ? ARGS[3] : "validation/outputs" mkpath(output_dir) model = CSV.read(model_file, DataFrame) expert = CSV.read(joinpath(run_dir, "expert_test.csv"), DataFrame) text = CSV.read(joinpath(run_dir, "text_data.csv"), DataFrame) party_col = :party_id in propertynames(model) ? :party_id : :party model_lookup = Dict{Tuple{Int,Int}, NamedTuple}() for row in eachrow(model) model_lookup[(Int(row[party_col]), Int(row.year))] = ( economic_lr = row.economic_lr, economic_lr_q025 = row.economic_lr_q025, economic_lr_q975 = row.economic_lr_q975, galtan = row.galtan, galtan_q025 = row.galtan_q025, galtan_q975 = row.galtan_q975) end text_years = Dict{Int, Vector{Int}}() for d in groupby(text, :party) text_years[Int(first(d.party))] = sort(unique(Int.(d.year))) end rows = NamedTuple[] for row in eachrow(expert) key = (Int(row.party), Int(row.year)) haskey(model_lookup, key) || continue dimension = row.var in ECON_VARS ? "economic_lr" : row.var in CULT_VARS ? "galtan" : nothing isnothing(dimension) && continue m = model_lookup[key] if dimension == "economic_lr" value, lower, upper = m.economic_lr, m.economic_lr_q025, m.economic_lr_q975 else value, lower, upper = m.galtan, m.galtan_q025, m.galtan_q975 end years = get(text_years, Int(row.party), Int[]) isempty(years) && continue distance = minimum(abs.(years .- Int(row.year))) distance_class = distance == 0 ? "direct text" : distance <= 3 ? "nearby text (1--3 years)" : "distant text (4+ years)" err = Float64(row.val) - value push!(rows, ( party_id = Int(row.party), country = String(row.country), region = region_for(String(row.country)), year = Int(row.year), decade = string(fld(Int(row.year), 10) * 10, "s"), project = String(row.project), item = String(row.var), dimension = dimension, expert_value = Float64(row.val), model_value = value, interval_lower = lower, interval_upper = upper, error = err, absolute_error = abs(err), within_latent_interval_95 = lower <= Float64(row.val) <= upper, nearest_text_distance = distance, text_distance_class = distance_class )) end predictions = DataFrame(rows) isempty(predictions) && error("No held-out expert observations matched the blocked-fit output") summaries = DataFrame[] for dimension in sort(unique(predictions.dimension)) d_dim = predictions[predictions.dimension .== dimension, :] push!(summaries, summarize_group(d_dim, "overall", "All matched held-out ratings")) for variable in [:decade, :region, :text_distance_class, :project, :item] for d_group in groupby(d_dim, variable) push!(summaries, summarize_group(d_group, String(variable), string(first(d_group[!, variable])))) end end end summary = vcat(summaries...) predictions_file = joinpath(output_dir, "blocked_validation_predictions.csv") summary_file = joinpath(output_dir, "blocked_validation_summary.csv") CSV.write(predictions_file, predictions) CSV.write(summary_file, summary) println("Wrote $predictions_file ($(nrow(predictions)) rows)") println("Wrote $summary_file ($(nrow(summary)) rows)") println(summary[summary.group_type .== "overall", :]) end main()