Files
party2d/src/julia/02_post_estimation.jl
T

987 lines
36 KiB
Julia

#!/usr/bin/env julia
#=
02_post_estimation.jl - Extract party position estimates from Stan model output
Supports both model versions:
- 2D model (V1): Extracts economic left-right and cultural cosmopolitan--traditionalist positions directly
- 4D model (V10): Extracts 4 traits + 2 derived scales
V10/V1 UPDATE: Handles segment-based indexing
- Maps segment results back to original party IDs
- Adds segment_num column to indicate which segment of the party
- Flags discontinuities for parties with multiple segments
This script:
1. Auto-detects the latest model run in model_outputs/
2. Loads the chain CSV files and data mappings
3. Detects model version from metadata or column names
4. Extracts posterior summaries for all segment-year positions
5. Maps Stan parameter indices back to real party IDs, segment numbers, and years
6. Saves output as wide-format CSV with uncertainty estimates
Usage:
julia 02_post_estimation.jl
Output:
party_positions_YYYY-MM-DD_HH-MM-SS.csv
=#
using CSV
using DataFrames
using Statistics
using JSON
using Dates
using Printf
# =============================================================================
# STEP 0: Auto-detect latest run
# =============================================================================
function find_latest_run(base_dir::String="outputs/model_outputs/latest")
if !isdir(base_dir)
error("Model outputs directory not found: $base_dir")
end
runs = filter(d -> startswith(d, "run_") && isdir(joinpath(base_dir, d)), readdir(base_dir))
if isempty(runs)
error("No runs found in $base_dir")
end
# Sort by timestamp in directory name (format: run_YYYY-MM-DD_HH-MM-SS)
sort!(runs, rev=true)
latest = joinpath(base_dir, runs[1])
println("Found $(length(runs)) run(s). Using latest: $latest")
return latest
end
# =============================================================================
# STEP 1: Load data and build segment-year lookup
# =============================================================================
function load_run_data(run_dir::String)
println("\n" * "="^60)
println("LOADING RUN DATA")
println("="^60)
data_dir = joinpath(run_dir, "data")
chains_dir = joinpath(run_dir, "chains")
# Check required files exist
required_files = [
joinpath(data_dir, "text_data.csv"),
joinpath(data_dir, "expert_dim.csv"),
joinpath(data_dir, "expert_lr.csv"),
joinpath(run_dir, "metadata.json")
]
for f in required_files
if !isfile(f)
error("Required file not found: $f")
end
end
# Load data files
println("Loading text_data.csv...")
text_data = CSV.read(joinpath(data_dir, "text_data.csv"), DataFrame)
println(" Rows: $(nrow(text_data))")
println("Loading expert_dim.csv...")
expert_dim = CSV.read(joinpath(data_dir, "expert_dim.csv"), DataFrame)
println(" Rows: $(nrow(expert_dim))")
println("Loading expert_lr.csv...")
expert_lr = CSV.read(joinpath(data_dir, "expert_lr.csv"), DataFrame)
println(" Rows: $(nrow(expert_lr))")
println("Loading metadata.json...")
metadata = JSON.parsefile(joinpath(run_dir, "metadata.json"))
println(" year0: $(metadata["year0"])")
println(" Model: $(metadata["model_file"])")
# V10: Load segment_info if available
segment_info_file = joinpath(data_dir, "segment_info.csv")
segment_info = nothing
if isfile(segment_info_file)
println("Loading segment_info.csv (V10)...")
segment_info = CSV.read(segment_info_file, DataFrame)
println(" Segments: $(nrow(segment_info))")
end
# V10: Load segment_year_map if available
segment_year_file = joinpath(data_dir, "segment_year_map.csv")
segment_year_map = nothing
if isfile(segment_year_file)
println("Loading segment_year_map.csv (V10)...")
segment_year_map = CSV.read(segment_year_file, DataFrame)
println(" Segment-years: $(nrow(segment_year_map))")
end
# Find chain files
chain_files = filter(f -> endswith(f, ".csv") && startswith(f, "chain_"), readdir(chains_dir))
println("\nFound $(length(chain_files)) chain file(s)")
return (
text_data = text_data,
expert_dim = expert_dim,
expert_lr = expert_lr,
metadata = metadata,
segment_info = segment_info,
segment_year_map = segment_year_map,
chain_files = [joinpath(chains_dir, f) for f in sort(chain_files)],
run_dir = run_dir
)
end
function normalize_country_value(value)
if ismissing(value)
return missing
end
txt = strip(string(value))
return isempty(txt) ? missing : txt
end
function build_party_country_map(text_data::DataFrame, expert_dim::DataFrame, expert_lr::DataFrame)
merged = unique(vcat(
select(text_data, :party, :country),
select(expert_dim, :party, :country),
select(expert_lr, :party, :country)
))
party_to_country = Dict{Int, String}()
for row in eachrow(merged)
pid = tryparse(Int, string(row.party))
if pid === nothing
continue
end
c = normalize_country_value(row.country)
if !ismissing(c)
party_to_country[pid] = c
end
end
return party_to_country
end
function load_constituent_to_union_map()::Dict{Int, Int}
mapping_file = joinpath("data", "union_mapping.csv")
constituent_to_union = Dict{Int, Int}()
if isfile(mapping_file)
union_df = CSV.read(mapping_file, DataFrame)
for row in eachrow(union_df)
constituent_to_union[row.expert_pf_id] = row.manifesto_pf_id
end
end
return constituent_to_union
end
function resolve_party_country(pid_value,
party_to_country::Dict{Int, String},
constituent_to_union::Dict{Int, Int})
pid = tryparse(Int, string(pid_value))
if pid === nothing
return missing, "unresolved"
end
if haskey(party_to_country, pid)
return party_to_country[pid], "direct"
end
if haskey(constituent_to_union, pid)
uid = constituent_to_union[pid]
if haskey(party_to_country, uid)
return party_to_country[uid], "union_fallback"
end
end
return missing, "unresolved"
end
function apply_country_resolution!(df::DataFrame,
party_col::Symbol,
country_col::Symbol,
party_to_country::Dict{Int, String},
constituent_to_union::Dict{Int, Int})
resolved_country = Union{Missing, String}[]
source_counts = Dict("direct" => 0, "union_fallback" => 0, "unresolved" => 0)
unresolved_parties = Set{Int}()
for pid in df[!, party_col]
country, source = resolve_party_country(pid, party_to_country, constituent_to_union)
push!(resolved_country, country)
source_counts[source] += 1
if source == "unresolved"
pid_int = tryparse(Int, string(pid))
if pid_int !== nothing
push!(unresolved_parties, pid_int)
end
end
end
df[!, country_col] = resolved_country
return source_counts, sort!(collect(unresolved_parties))
end
function fill_missing_countries!(df::DataFrame,
segment_info::Union{DataFrame, Nothing},
party_to_country::Dict{Int, String},
constituent_to_union::Dict{Int, Int})
if !hasproperty(df, :country)
source_counts, unresolved = apply_country_resolution!(
df, :party_id, :country, party_to_country, constituent_to_union
)
return source_counts, unresolved
end
normalized = Union{Missing, String}[]
for val in df.country
push!(normalized, normalize_country_value(val))
end
df.country = normalized
source_counts = Dict("direct" => 0, "union_fallback" => 0, "segment_info" => 0, "unresolved" => 0)
segment_country_by_id = Dict{Int, String}()
if segment_info !== nothing && hasproperty(segment_info, :country)
for row in eachrow(segment_info)
c = normalize_country_value(row.country)
if !ismissing(c)
segment_country_by_id[Int(row.segment_id)] = c
end
end
end
unresolved_parties = Set{Int}()
for i in 1:nrow(df)
if !ismissing(df.country[i])
continue
end
if hasproperty(df, :segment_id) && haskey(segment_country_by_id, Int(df.segment_id[i]))
df.country[i] = segment_country_by_id[Int(df.segment_id[i])]
source_counts["segment_info"] += 1
continue
end
country, source = resolve_party_country(df.party_id[i], party_to_country, constituent_to_union)
if !ismissing(country)
df.country[i] = country
source_counts[source] += 1
else
source_counts["unresolved"] += 1
pid_int = tryparse(Int, string(df.party_id[i]))
pid_int !== nothing && push!(unresolved_parties, pid_int)
end
end
return source_counts, sort!(collect(unresolved_parties))
end
# =============================================================================
# STEP 2: Build complete segment-year mapping (V10) or party-year mapping (V9)
# =============================================================================
function build_segment_year_map(text_data::DataFrame, expert_dim::DataFrame, expert_lr::DataFrame,
segment_info::Union{DataFrame, Nothing},
segment_year_map::Union{DataFrame, Nothing},
run_dir::String,
year0::Int)
println("\n" * "="^60)
println("BUILDING SEGMENT-YEAR MAPPING")
println("="^60)
party_to_country = build_party_country_map(text_data, expert_dim, expert_lr)
constituent_to_union = load_constituent_to_union_map()
# V10: Use segment_year_map if available
if segment_year_map !== nothing && segment_info !== nothing
println("Using segment_year_map.csv (V10 mode)")
# Convert relative Year to absolute year
if hasproperty(segment_year_map, :Year)
segment_year_map.year = segment_year_map.Year .+ year0
elseif !hasproperty(segment_year_map, :year)
error("segment_year_map has no Year or year column")
end
# Add party_id from segment_info if not already present
if !hasproperty(segment_year_map, :party_id)
segment_id_to_party = Dict(row.segment_id => row.party_id for row in eachrow(segment_info))
segment_year_map.party_id = [segment_id_to_party[sid] for sid in segment_year_map.segment_id]
end
# Add segment_num from segment_info if not already present
if !hasproperty(segment_year_map, :segment_num)
segment_id_to_segnum = Dict(row.segment_id => row.segment_num for row in eachrow(segment_info))
segment_year_map.segment_num = [segment_id_to_segnum[sid] for sid in segment_year_map.segment_id]
end
# Resolve/fill country column using segment metadata first, then direct and union-fallback lookup.
source_counts, unresolved = fill_missing_countries!(
segment_year_map, segment_info, party_to_country, constituent_to_union
)
direct_count = get(source_counts, "direct", 0)
union_count = get(source_counts, "union_fallback", 0)
segment_info_count = get(source_counts, "segment_info", 0)
unresolved_count = count(ismissing, segment_year_map.country)
println(" Country resolution fill counts: direct=$direct_count, union_fallback=$union_count, segment_info=$segment_info_count, unresolved_rows=$unresolved_count")
if !isempty(unresolved)
println(" Warning: unresolved country party IDs (first 20): $(unresolved[1:min(20, length(unresolved))])")
end
R = maximum(segment_year_map.rr)
n_segments = length(unique(segment_year_map.segment_id))
n_parties = length(unique(segment_year_map.party_id))
println("Loaded segment_year_map: $(nrow(segment_year_map)) segment-years (R=$R)")
println(" Unique segments: $n_segments")
println(" Unique parties: $n_parties")
# Count observed vs interpolated
observed_rrs = Set{Int}()
if hasproperty(text_data, :rr_man)
union!(observed_rrs, Set(text_data.rr_man))
end
if hasproperty(expert_dim, :rr_exp_dim)
union!(observed_rrs, Set(expert_dim.rr_exp_dim))
end
if hasproperty(expert_lr, :rr_exp_lr)
union!(observed_rrs, Set(expert_lr.rr_exp_lr))
end
n_observed = length(intersect(Set(segment_year_map.rr), observed_rrs))
n_interpolated = nrow(segment_year_map) - n_observed
println(" Observed segment-years: $n_observed")
println(" Interpolated segment-years: $n_interpolated")
return segment_year_map, R, segment_info
end
# V9 fallback: Use party_year_map
party_year_file = joinpath(run_dir, "data", "party_year_map.csv")
if isfile(party_year_file)
println("Loading party_year_map.csv (V9 fallback mode)")
party_year_map = CSV.read(party_year_file, DataFrame)
# Add party_id column (same as party for V9)
if !hasproperty(party_year_map, :party_id)
party_year_map.party_id = party_year_map.party
end
# Add segment_num column (always 1 for V9)
if !hasproperty(party_year_map, :segment_num)
party_year_map.segment_num = ones(Int, nrow(party_year_map))
end
# Add segment_id column (same as party index for V9)
if !hasproperty(party_year_map, :segment_id)
party_year_map.segment_id = party_year_map.party
end
# Convert relative Year to absolute year
if hasproperty(party_year_map, :Year)
party_year_map.year = party_year_map.Year .+ year0
elseif !hasproperty(party_year_map, :year)
error("party_year_map has no Year or year column")
end
# Resolve/fill country column
if !hasproperty(party_year_map, :country)
source_counts, unresolved = apply_country_resolution!(
party_year_map, :party_id, :country, party_to_country, constituent_to_union
)
direct_count = source_counts["direct"]
union_count = source_counts["union_fallback"]
unresolved_count = source_counts["unresolved"]
println(" Country resolution sources: direct=$direct_count, union_fallback=$union_count, unresolved=$unresolved_count")
if !isempty(unresolved)
println(" Warning: unresolved country party IDs (first 20): $(unresolved[1:min(20, length(unresolved))])")
end
else
normalized = Union{Missing, String}[]
for val in party_year_map.country
push!(normalized, normalize_country_value(val))
end
party_year_map.country = normalized
end
R = maximum(party_year_map.rr)
println("Loaded party_year_map: $(nrow(party_year_map)) party-years (R=$R)")
return party_year_map, R, nothing
end
# Fallback: Reconstruct from data files
@warn "No mapping file found, reconstructing from data (observed years only)"
# Extract unique party-year-rr combinations from text_data
text_map = unique(select(text_data, :party, :country, :year, :rr_man))
rename!(text_map, :rr_man => :rr)
text_map.party_id = text_map.party
text_map.segment_num = ones(Int, nrow(text_map))
expert_dim_map = unique(select(expert_dim, :party, :country, :year, :rr_exp_dim))
rename!(expert_dim_map, :rr_exp_dim => :rr)
expert_dim_map.party_id = expert_dim_map.party
expert_dim_map.segment_num = ones(Int, nrow(expert_dim_map))
expert_lr_map = unique(select(expert_lr, :party, :country, :year, :rr_exp_lr))
rename!(expert_lr_map, :rr_exp_lr => :rr)
expert_lr_map.party_id = expert_lr_map.party
expert_lr_map.segment_num = ones(Int, nrow(expert_lr_map))
combined = vcat(text_map, expert_dim_map, expert_lr_map)
segment_year_map = unique(combined)
sort!(segment_year_map, :rr)
R = maximum(segment_year_map.rr)
println("Reconstructed mapping: $(nrow(segment_year_map)) segment-years (R=$R)")
return segment_year_map, R, nothing
end
# =============================================================================
# STEP 3: Load and combine chains
# =============================================================================
function load_chains(chain_files::Vector{String})
println("\n" * "="^60)
println("LOADING STAN CHAINS")
println("="^60)
flush(stdout)
chains = DataFrame[]
# The full Stan CSVs are very wide (hundreds of thousands of columns). For
# post-estimation we only need party-position generated quantities. Reading
# all columns can take hours and allocate many GB of irrelevant parameters.
post_estimation_prefixes = (
"economic_lr.",
"galtan.",
"pro_market.",
"pro_welfare.",
"cosmopolitan.",
"traditional.",
)
keep_post_estimation_col(_i, name) = any(startswith(String(name), p) for p in post_estimation_prefixes)
for (i, f) in enumerate(chain_files)
println("Loading chain $i: $(basename(f))...")
flush(stdout)
# Skip comment lines (Stan header) and parse only needed quantities.
chain = CSV.read(f, DataFrame; comment="#", select=keep_post_estimation_col)
println(" Samples: $(nrow(chain)), Parameters: $(ncol(chain))")
flush(stdout)
push!(chains, chain)
end
# Combine chains
println("Combining selected chain columns...")
flush(stdout)
combined = vcat(chains...)
println("\nCombined: $(nrow(combined)) total samples")
println("Selected parameters: $(ncol(combined))")
flush(stdout)
return combined
end
# =============================================================================
# STEP 4: Extract generated quantities
# =============================================================================
"""
Detect model version from chain column names.
Returns "2dim" or "4dim".
"""
function detect_model_version(chains::DataFrame)
cols = names(chains)
# 2D model has economic_lr but NOT pro_market
has_economic_lr = any(c -> startswith(string(c), "economic_lr."), cols)
has_pro_market = any(c -> startswith(string(c), "pro_market."), cols)
if has_economic_lr && !has_pro_market
return "2dim"
elseif has_pro_market
return "4dim"
else
error("Could not detect model version from chain columns")
end
end
function extract_estimates(chains::DataFrame, segment_year_map::DataFrame, R::Int)
println("\n" * "="^60)
println("EXTRACTING POSTERIOR ESTIMATES")
println("="^60)
# Auto-detect model version from columns
model_version = detect_model_version(chains)
println("Detected model version: $model_version")
# Select quantities based on model version
if model_version == "2dim"
# 2D model: economic left-right and cultural cosmopolitan--traditionalist positions are directly estimated
# (general_lr is computed in Stan for anchoring but not extracted as output)
quantities = ["economic_lr", "galtan"]
test_col = "economic_lr.1"
else
# 4D model: 4 traits + 2 derived scales
quantities = ["pro_market", "pro_welfare", "cosmopolitan", "traditional", "economic_lr", "galtan"]
test_col = "pro_market.1"
end
# Check that columns exist
if !hasproperty(chains, Symbol(test_col))
error("Column $test_col not found in chains. Available columns: $(first(names(chains), 10))...")
end
n_samples = nrow(chains)
println("Samples per parameter: $n_samples")
# Load union mapping for adding union_party_id column
union_mapping_file = joinpath("data", "union_mapping.csv")
constituent_to_union_pf = Dict{Int, Int}()
if isfile(union_mapping_file)
union_df = CSV.read(union_mapping_file, DataFrame)
for row in eachrow(union_df)
constituent_to_union_pf[row.expert_pf_id] = row.manifesto_pf_id
end
end
# Pre-allocate output DataFrame
n_rows = nrow(segment_year_map)
# Add union_party_id column: NA for standalone parties, union PF ID for constituents
union_ids = Union{Int, Missing}[]
for pid in segment_year_map.party_id
pid_int = isa(pid, Integer) ? pid : tryparse(Int, string(pid))
if pid_int !== nothing && haskey(constituent_to_union_pf, pid_int)
push!(union_ids, constituent_to_union_pf[pid_int])
else
push!(union_ids, missing)
end
end
output = DataFrame(
party_id = segment_year_map.party_id,
union_party_id = union_ids,
segment_num = segment_year_map.segment_num,
country = segment_year_map.country,
year = segment_year_map.year,
rr = segment_year_map.rr
)
# Add columns for each quantity
for q in quantities
output[!, Symbol(q)] = zeros(Float64, n_rows)
output[!, Symbol("$(q)_se")] = zeros(Float64, n_rows)
output[!, Symbol("$(q)_q025")] = zeros(Float64, n_rows)
output[!, Symbol("$(q)_q975")] = zeros(Float64, n_rows)
end
println("Extracting estimates for $(n_rows) segment-year positions...")
# Progress tracking
prog_interval = max(1, n_rows ÷ 20)
for (i, row) in enumerate(eachrow(segment_year_map))
r = row.rr
# Progress
if i % prog_interval == 0 || i == n_rows
pct = round(100 * i / n_rows, digits=1)
print("\r Progress: $pct% ($i / $n_rows)")
end
for q in quantities
col_name = Symbol("$q.$r")
if !hasproperty(chains, col_name)
@warn "Column $col_name not found (rr=$r)" maxlog=5
continue
end
samples = chains[!, col_name]
# Compute summary statistics
output[i, Symbol(q)] = mean(samples)
output[i, Symbol("$(q)_se")] = std(samples)
output[i, Symbol("$(q)_q025")] = quantile(samples, 0.025)
output[i, Symbol("$(q)_q975")] = quantile(samples, 0.975)
end
end
println() # Newline after progress
# Remove the rr column from final output (internal only)
select!(output, Not(:rr))
return output
end
# =============================================================================
# STEP 5: Validation
# =============================================================================
function validate_output(output::DataFrame, segment_info::Union{DataFrame, Nothing})
println("\n" * "="^60)
println("VALIDATION CHECKS")
println("="^60)
all_passed = true
# Detect which columns are present (2D vs 4D model)
has_4d = hasproperty(output, :pro_market)
# Check 1: Range check - all estimates should be in [0, 1]
println("\n1. Range check (all values in [0, 1]):")
if has_4d
check_cols = [:pro_market, :pro_welfare, :cosmopolitan, :traditional, :economic_lr, :galtan]
else
check_cols = [:economic_lr, :galtan]
end
for col in check_cols
if !hasproperty(output, col)
continue
end
vals = output[!, col]
min_val, max_val = extrema(vals)
in_range = min_val >= 0 && max_val <= 1
status = in_range ? "PASS" : "FAIL"
println(" $col: [$(@sprintf("%.4f", min_val)), $(@sprintf("%.4f", max_val))] - $status")
all_passed = all_passed && in_range
end
# Check 2: Anchor party checks
println("\n2. Anchor party checks:")
# Define anchor parties with expected ranges (for 2D model)
# Includes both union IDs (V3) and individual constituent IDs (V4)
anchor_parties = [
(id=211, name="CDU/CSU", country="DE", econ=(0.50, 0.70), galtan=(0.45, 0.70)),
(id=1375, name="CDU", country="DE", econ=(0.50, 0.70), galtan=(0.45, 0.65)),
(id=1731, name="CSU", country="DE", econ=(0.50, 0.70), galtan=(0.55, 0.75)),
(id=383, name="SPD", country="DE", econ=(0.30, 0.50), galtan=(0.30, 0.55)),
(id=1516, name="Labour", country="GB", econ=(0.30, 0.55), galtan=(0.30, 0.55)),
(id=1567, name="Conservatives", country="GB", econ=(0.55, 0.80), galtan=(0.50, 0.75)),
(id=487, name="SAP", country="SE", econ=(0.30, 0.50), galtan=(0.35, 0.55)),
]
n_checked = 0
n_passed = 0
for anchor in anchor_parties
party_rows = filter(r -> r.party_id == anchor.id, output)
if nrow(party_rows) == 0
println(" $(anchor.name) ($(anchor.id)): NOT FOUND")
continue
end
# Use most recent 20 years of data as reference period
max_year = maximum(party_rows.year)
ref_rows = filter(r -> r.year >= max_year - 20, party_rows)
if nrow(ref_rows) == 0
ref_rows = party_rows
end
n_checked += 1
mean_econ = mean(ref_rows.economic_lr)
mean_galtan = mean(ref_rows.galtan)
econ_ok = anchor.econ[1] <= mean_econ <= anchor.econ[2]
galtan_ok = anchor.galtan[1] <= mean_galtan <= anchor.galtan[2]
all_ok = econ_ok && galtan_ok
if all_ok
n_passed += 1
end
status = all_ok ? "PASS" : "WARN"
econ_marker = econ_ok ? "" : "*"
galtan_marker = galtan_ok ? "" : "*"
@printf(" %-15s economic=%.2f%s [%.2f-%.2f] cultural=%.2f%s [%.2f-%.2f] %s\n",
anchor.name, mean_econ, econ_marker, anchor.econ[1], anchor.econ[2],
mean_galtan, galtan_marker, anchor.galtan[1], anchor.galtan[2], status)
end
if n_checked > 0
println(" Anchor check: $n_passed/$n_checked within expected ranges")
println(" Note: Model integrates text + expert data; deviations from expert-only expectations are normal")
end
# Check 3: Coverage check
println("\n3. Coverage check:")
println(" Total segment-year positions: $(nrow(output))")
println(" Unique parties: $(length(unique(output.party_id)))")
blank_country_rows = count(ismissing, output.country)
println(" Unique countries: $(length(unique(skipmissing(output.country))))")
if blank_country_rows == 0
println(" Blank country rows: 0 - PASS")
else
println(" Blank country rows: $blank_country_rows - FAIL")
all_passed = false
end
println(" Year range: $(minimum(output.year)) - $(maximum(output.year))")
# V10: Check segment distribution
segment_counts = combine(groupby(output, :party_id), nrow => :n_years,
:segment_num => (x -> length(unique(x))) => :n_segments)
parties_multi_segment = filter(:n_segments => >(1), segment_counts)
if nrow(parties_multi_segment) > 0
println("\n Parties with multiple segments: $(length(unique(parties_multi_segment.party_id)))")
end
# Check 4: No duplicates (party_id, segment_num, year should be unique)
println("\n4. Duplicate check:")
dup_count = nrow(output) - nrow(unique(select(output, :party_id, :segment_num, :year)))
if dup_count == 0
println(" No duplicate (party_id, segment_num, year) combinations - PASS")
else
println(" WARNING: Found $dup_count duplicate combinations!")
all_passed = false
end
# Check 5: SE reasonableness
println("\n5. Standard error check:")
se_cols = has_4d ?
[:pro_market_se, :pro_welfare_se, :cosmopolitan_se, :traditional_se] :
[:economic_lr_se, :galtan_se]
for col in se_cols
if !hasproperty(output, col)
continue
end
vals = output[!, col]
mean_se = mean(vals)
max_se = maximum(vals)
println(" $col: mean=$(@sprintf("%.4f", mean_se)), max=$(@sprintf("%.4f", max_se))")
end
println("\n" * "-"^60)
if all_passed
println("All validation checks PASSED")
else
println("Some validation checks FAILED - please inspect output carefully")
end
return all_passed
end
# =============================================================================
# STEP 6: Save output
# =============================================================================
function save_output(output::DataFrame, metadata::Dict, segment_info::Union{DataFrame, Nothing}, run_dir::String; outdir::String="outputs/estimations/latest")
println("\n" * "="^60)
println("SAVING OUTPUT")
println("="^60)
timestamp = Dates.format(now(), "yyyy-mm-dd_HH-MM-SS")
mkpath(outdir)
# Delete previous output files
for f in readdir(outdir)
if startswith(f, "party_positions_") && (endswith(f, ".csv") || endswith(f, ".txt") || endswith(f, ".tex"))
rm(joinpath(outdir, f))
println(" Deleted old: $f")
end
end
# Save main CSV
csv_file = joinpath(outdir, "party_positions_$timestamp.csv")
CSV.write(csv_file, output)
println("Saved: $csv_file")
println(" Rows: $(nrow(output))")
println(" Columns: $(ncol(output))")
# Count parties with multiple segments
n_multi_segment = 0
if segment_info !== nothing
party_segment_counts = combine(groupby(segment_info, :party_id), nrow => :n_segments)
n_multi_segment = count(r -> r.n_segments > 1, eachrow(party_segment_counts))
end
# Save metadata
meta_file = joinpath(outdir, "party_positions_$(timestamp)_metadata.txt")
open(meta_file, "w") do f
println(f, "Party Positions Dataset - Metadata")
println(f, "="^50)
println(f, "")
println(f, "Generated: $(Dates.format(now(), "yyyy-mm-dd HH:MM:SS"))")
println(f, "Source run: $(basename(run_dir))")
println(f, "Model file: $(get(metadata, "model_file", "unknown"))")
println(f, "")
println(f, "Dataset size:")
println(f, " Segment-year observations: $(nrow(output))")
println(f, " Unique parties: $(length(unique(output.party_id)))")
if n_multi_segment > 0
println(f, " Parties with multiple segments: $n_multi_segment")
end
println(f, " Unique countries: $(length(unique(output.country)))")
println(f, " Year range: $(minimum(output.year)) - $(maximum(output.year))")
println(f, "")
println(f, "Columns:")
println(f, " party_id: PartyFacts ID (integer) - individual party (e.g., CDU=1375, CSU=1731)")
println(f, " union_party_id: PartyFacts ID of parent union (NA for standalone parties)")
println(f, " segment_num: Segment number within party (1, 2, 3...)")
println(f, " country: ISO2 country code")
println(f, " year: Calendar year")
println(f, "")
println(f, "Segment-Based Indexing:")
println(f, " - Parties are split into segments at gaps > 7 years")
println(f, " - Each segment is estimated independently (no continuity across gaps)")
println(f, " - Segments with < 3 observations are dropped")
println(f, " - segment_num=1 is the main segment; higher numbers indicate gaps in data")
println(f, "")
# Check if this is 2D or 4D output
is_2d = !hasproperty(output, :pro_market)
if is_2d
println(f, "Model: 2D Direct Bipolar")
println(f, "")
println(f, "Bipolar scales (0 = left/cosmopolitan, 1 = right/traditionalist):")
println(f, " economic_lr: Economic left-right position (directly estimated)")
println(f, " galtan: Cultural cosmopolitan--traditionalist position (directly estimated)")
# Note: general_lr is computed internally for cross-dimensional anchoring
# but not reported as output (the two dimension-specific estimates are preferred)
else
println(f, "Model: 4D Unipolar")
println(f, "")
println(f, "Dimensions (0 = low, 1 = high):")
println(f, " pro_market: Pro-market economic position")
println(f, " pro_welfare: Pro-welfare state position")
println(f, " cosmopolitan: Cosmopolitan cultural position")
println(f, " traditional: Traditionalist cultural position")
println(f, "")
println(f, "Derived bipolar scales (0 = left/cosmopolitan, 1 = right/traditionalist):")
println(f, " economic_lr: Economic left-right (derived from pro_market - pro_welfare)")
println(f, " galtan: Cultural cosmopolitan--traditionalist (derived from traditional - cosmopolitan)")
end
println(f, "")
println(f, "Uncertainty columns:")
println(f, " *_se: Standard error (posterior SD)")
println(f, " *_q025: 2.5th percentile (lower 95% CI)")
println(f, " *_q975: 97.5th percentile (upper 95% CI)")
println(f, "")
println(f, "Model convergence:")
println(f, " Mean R-hat: $(get(metadata, "mean_rhat", "N/A"))")
println(f, " Max R-hat: $(get(metadata, "max_rhat", "N/A"))")
println(f, " Mean ESS: $(get(metadata, "mean_ess", "N/A"))")
println(f, " Min ESS: $(get(metadata, "min_ess", "N/A"))")
end
println("Saved: $meta_file")
return csv_file, meta_file
end
# =============================================================================
# STEP 5b: Verify no union/alliance IDs in output
# =============================================================================
function verify_no_unions_in_output(output::DataFrame)
println("\n" * "="^60)
println("UNION ID VERIFICATION")
println("="^60)
union_mapping_file = joinpath("data", "union_mapping.csv")
if !isfile(union_mapping_file)
println(" No union_mapping.csv found — skipping verification")
return
end
union_df = CSV.read(union_mapping_file, DataFrame)
union_pf_ids = Set(union_df.manifesto_pf_id)
output_pf_ids = Set(output.party_id)
violations = intersect(union_pf_ids, output_pf_ids)
if isempty(violations)
println(" PASS: No union/alliance PF IDs found in output")
println(" Checked $(length(union_pf_ids)) union IDs against $(length(output_pf_ids)) output parties")
else
println(" WARNING: $(length(violations)) union PF IDs found in output")
println(" (This is expected if union_mapping.csv was updated after the model run)")
for v in sort(collect(violations))
n_rows = count(r -> r.party_id == v, eachrow(output))
println(" PF $v: $n_rows rows")
end
end
end
# =============================================================================
# MAIN
# =============================================================================
function main()
println("="^60)
println("POST-ESTIMATION: Party-position model")
println("="^60)
println("Started: $(Dates.format(now(), "yyyy-mm-dd HH:MM:SS"))")
# Step 0: Find run directory (CLI --run-dir or auto-detect latest)
run_dir = nothing
output_dir = nothing
for (i, arg) in enumerate(ARGS)
if arg == "--run-dir" && i < length(ARGS)
run_dir = ARGS[i + 1]
elseif startswith(arg, "--run-dir=")
run_dir = split(arg, "=", limit=2)[2]
elseif arg == "--output-dir" && i < length(ARGS)
output_dir = ARGS[i + 1]
elseif startswith(arg, "--output-dir=")
output_dir = split(arg, "=", limit=2)[2]
end
end
if run_dir === nothing
run_dir = find_latest_run()
else
println("Using specified run directory: $run_dir")
end
# Step 1: Load run data
data = load_run_data(run_dir)
# Step 2: Build segment-year mapping
year0 = data.metadata["year0"]
segment_year_map, R, segment_info = build_segment_year_map(
data.text_data, data.expert_dim, data.expert_lr,
data.segment_info, data.segment_year_map, data.run_dir, year0
)
# Step 3: Load chains
chains = load_chains(data.chain_files)
# Step 4: Extract estimates
output = extract_estimates(chains, segment_year_map, R)
# Step 5: Validate
validate_output(output, segment_info)
# Step 5b: Verify no union/alliance IDs in output
verify_no_unions_in_output(output)
# Step 6: Save output
effective_output_dir = output_dir !== nothing ? output_dir : "outputs/estimations/latest"
csv_file, meta_file = save_output(output, data.metadata, segment_info, run_dir; outdir=effective_output_dir)
println("\n" * "="^60)
println("COMPLETE")
println("="^60)
println("Output files:")
println(" $csv_file")
println(" $meta_file")
println("\nFinished: $(Dates.format(now(), "yyyy-mm-dd HH:MM:SS"))")
return output
end
# Run if executed directly
if abspath(PROGRAM_FILE) == @__FILE__
main()
end