Skip to content

Plotting

What Is It

survey_kit.plot builds interactive plotly figures directly from a StatCalculator's or MultipleImputation's own estimates - no manual reshaping. combine puts several figures on one page, switched between with dropdowns.

Key Features

  • No manual reshaping - point a plot function at a StatCalculator/MultipleImputation/AdapterStats and it finds the right columns itself.
  • DRB rounding by default - values are rounded per Census disclosure-review rules before plotting.
  • A group dropdown, not just a legend - line()/quantiles()/coefplot() can group series into a compact dropdown (group_by) instead of a long legend.
  • Confidence intervals built in - ci_level= draws error bars or, with ci_area=True, a shaded band.
  • combine() - nests any number of dropdown levels, remembering each level's last choice. Label each level (label=["Run:", "CI:"]), or put a level on its own row with a "\n" in its label.
  • Shared legend state - isolating a series by name on one figure carries over to any other figure with a same-named series, once it's shown.
  • Still a plain plotly figure - every function takes a layout: dict applied via fig.update_layout(**layout), and returns something you can call any plotly method on.

When to Use What

Use Case Tool
Any set of columns plotted against each other line()
A stat item's own quantile columns, percentile 0-100 quantiles()
One row per category with a CI whisker (disclosure-review style) coefplot()
Several stat items that should sum to a shown total stacked_bar()
Switching between whole figures (runs, vintages, scenarios) combine()

API

See the Plotting API reference for the full parameter list of every function.

Example/Tutorial

quantiles() for a stat item's own quantile columns, line() for any other set of columns. Covers confidence intervals and the group dropdown (group_by).

Figures: quantiles() · with CI · grouped · line()

import os
import polars as pl

from survey_kit.utilities.random import RandomData, set_seed, generate_seed
from survey_kit.statistics.calculator import StatCalculator
from survey_kit.statistics.statistics import Statistics
from survey_kit.statistics.replicates import Replicates
from survey_kit.statistics.bootstrap import bayes_bootstrap
from survey_kit import logger, config, plot

path_docs_plot = os.path.join(config.code_root, "..", "..", "docs", "tutorials", "plot")
path_docs_figures = os.path.join(path_docs_plot, "figures")
os.makedirs(path_docs_figures, exist_ok=True)

# %%
logger.info(
    "survey_kit.plot builds plotly figures directly from a StatCalculator's or "
    "MultipleImputation's own df_estimates/_df_ci - there's no separate "
    "reshaping step to do yourself. line()/quantiles() plot a set of stat "
    "columns (e.g. several quantiles) against each other, one line per index "
    "value or by-group."
)

set_seed(20260915)
n_rows = 2_000
n_replicates = 10

df = (
    RandomData(n_rows=n_rows, seed=generate_seed())
    .index("index")
    .integer("income", 0, 100_000)
    .integer("year", 2016, 2018)
).to_df()

df = pl.concat(
    [
        df,
        bayes_bootstrap(
            n_rows=n_rows,
            n_draws=n_replicates + 1,
            seed=generate_seed(),
            initial_weight_index=0,
            prefix="weight_",
        ),
    ],
    how="horizontal",
)

stats = Statistics(stats=["q10", "q25", "q50", "q75", "q90"], columns=["income"])
replicates = Replicates(
    weight_stub="weight_", n_replicates=n_replicates, bootstrap=True
)

sc = StatCalculator(
    df,
    statistics=stats,
    weight="weight_0",
    replicates=replicates,
    by={"year": ["year"]},
)
sc.print()

# %%
logger.info(
    "quantiles() is a thin wrapper around line() - it finds every quantile-stat "
    "column on its own (q10, q25, ...) and plots them across percentile 0-100. "
    "With more than one index column (here: Variable + year), each unique "
    "combination becomes its own line - one per year, since there's only one "
    "Variable (income) in this example."
)
fig_quantiles = plot.quantiles(sc)
fig_quantiles.write_html(
    os.path.join(path_docs_figures, "quantiles.html"), include_plotlyjs="directory"
)

# %%
logger.info(
    "The legend itself is always interactive, whether or not a figure has a "
    "group dropdown (see below): single-click a legend entry to toggle just "
    "that line on/off; double-click one to isolate it, hiding every other "
    "line at once (double-click again, or on another entry, to bring the "
    "rest back)."
)

# %%
logger.info(
    "Add confidence intervals with ci_level - either as error bars (the "
    "default) or, with ci_area=True, a shaded band around each line."
)
fig_quantiles_ci = plot.quantiles(sc, ci_level=0.95, ci_area=True)
fig_quantiles_ci.write_html(
    os.path.join(path_docs_figures, "quantiles_ci_area.html"), include_plotlyjs="directory"
)

# %%
logger.info(
    "Every line()/quantiles() figure carries a dropdown for switching which "
    "group of lines is shown - by default there's a single group (everything "
    "shown at once, same as a plain legend, and no extra JavaScript at all, "
    "since there's nothing to switch between). group_by splits each line's "
    "label on a separator and groups lines that share the same prefix - handy "
    "when several related series (e.g. two 'historical' years vs. a 'recent' "
    "one) should travel together in the dropdown instead of each getting its "
    "own entry."
)
fig_grouped = plot.quantiles(
    sc,
    rename={
        "income / 2016.0": "Historical:2016",
        "income / 2017.0": "Historical:2017",
        "income / 2018.0": "Recent:2018",
    },
    group_by=":",
    group_first=True,
)
logger.info(f"Groups: {fig_grouped._group_map}")
fig_grouped.write_html(
    os.path.join(path_docs_figures, "quantiles_grouped.html"), include_plotlyjs="directory"
)

# %%
logger.info(
    "line() itself is more general than quantiles() - point it at any set of "
    "df_estimates columns you want plotted against each other. Here we only "
    "plot 3 of the 5 quantile columns, in an explicit order (also filters to "
    "just those series)."
)
fig_line = plot.line(
    sc,
    x_columns=["q10", "q50", "q90"],
    x_axis_title="Percentile",
    y_axis_title="Income",
    order=["income / 2018.0", "income / 2017.0", "income / 2016.0"],
)
fig_line.write_html(
    os.path.join(path_docs_figures, "line_selected_quantiles.html"),
    include_plotlyjs="directory",
)

A disclosure-review-style point-and-whisker chart, with and without series= for multiple offset points per row.

Figures: with series · single series

import os
import narwhals as nw
import polars as pl

from survey_kit.utilities.random import RandomData, set_seed, generate_seed
from survey_kit.statistics.calculator import StatCalculator
from survey_kit.statistics.statistics import Statistics
from survey_kit.statistics.replicates import Replicates
from survey_kit.statistics.bootstrap import bayes_bootstrap
from survey_kit import logger, config, plot

path_docs_plot = os.path.join(config.code_root, "..", "..", "docs", "tutorials", "plot")
path_docs_figures = os.path.join(path_docs_plot, "figures")
os.makedirs(path_docs_figures, exist_ok=True)

# %%
logger.info(
    "coefplot() is the classic disclosure-review chart: one row per category "
    "(e.g. a demographic subgroup), a point with a confidence interval whisker "
    "for the estimate, and (optionally) several offset points per row when "
    "there's more than one series to compare (e.g. one color per year)."
)

set_seed(20260915)
n_rows = 3_000
n_replicates = 8
labels = ["Male", "Female", "White", "Black", "Under 18", "65+"]

df = (
    RandomData(n_rows=n_rows, seed=generate_seed())
    .index("index")
    .integer("poverty", 0, 1)
    .integer("subgroup_id", 0, len(labels) - 1)
    .integer("year", 2016, 2018)
).to_df()

when = pl
for i, label in enumerate(labels):
    when = when.when(pl.col("subgroup_id") == i).then(pl.lit(label))
df = df.with_columns(when.otherwise(pl.lit("Other")).alias("subgroup")).drop(
    "subgroup_id"
)

df = pl.concat(
    [
        df,
        bayes_bootstrap(
            n_rows=n_rows,
            n_draws=n_replicates + 1,
            seed=generate_seed(),
            initial_weight_index=0,
            prefix="weight_",
        ),
    ],
    how="horizontal",
)

stats = Statistics(stats=["mean"], columns=["poverty"])
replicates = Replicates(
    weight_stub="weight_", n_replicates=n_replicates, bootstrap=True
)

sc = StatCalculator(
    df,
    statistics=stats,
    weight="weight_0",
    replicates=replicates,
    by={"group": ["subgroup", "year"]},
)
sc.print()

# %%
logger.info(
    "category is the row axis (subgroup here), series offsets multiple points "
    "per row (year). headers insert a bold divider row above a given "
    "category - handy for grouping related subgroups (e.g. every "
    "gender/race/age category under one section heading)."
)
fig_coefplot = plot.coefplot(
    sc,
    column="mean",
    ci_level=0.95,
    category="subgroup",
    series="year",
    headers={"Male": "Gender", "White": "Race", "Under 18": "Age"},
    x_axis_title="Poverty rate",
)
fig_coefplot.write_html(
    os.path.join(path_docs_figures, "coefplot.html"), include_plotlyjs="directory"
)

# %%
logger.info(
    "Without `series`, coefplot() draws a single unoffset point per category - "
    "useful for a simple one-estimate-per-row chart (e.g. just the most "
    "recent year)."
)
sc_2018 = sc.filter(nw.col("year") == 2018)
fig_coefplot_single = plot.coefplot(
    sc_2018,
    column="mean",
    ci_level=0.95,
    category="subgroup",
    headers={"Male": "Gender", "White": "Race", "Under 18": "Age"},
    x_axis_title="Poverty rate, 2018",
)
fig_coefplot_single.write_html(
    os.path.join(path_docs_figures, "coefplot_single_series.html"),
    include_plotlyjs="directory",
)

Layers/categories can mix positive and negative values - each bar stacks positive layers right of zero and negative ones left, with the total label landing on whichever side its net value falls on.

Figures: all-negative · mixed sign

import os
import polars as pl

from survey_kit.utilities.random import RandomData, set_seed, generate_seed
from survey_kit.statistics.calculator import StatCalculator
from survey_kit.statistics.statistics import Statistics
from survey_kit.statistics.replicates import Replicates
from survey_kit.statistics.bootstrap import bayes_bootstrap
from survey_kit import logger, config, plot

path_docs_plot = os.path.join(config.code_root, "..", "..", "docs", "tutorials", "plot")
path_docs_figures = os.path.join(path_docs_plot, "figures")
os.makedirs(path_docs_figures, exist_ok=True)

# %%
logger.info(
    "stacked_bar() answers a different question than line()/coefplot(): given "
    "several StatCalculator/MultipleImputation objects that all share the same "
    "category axis (e.g. 'which safety-net program was removed'), how do their "
    "contributions add up? Each dict entry becomes one layer of the stack (e.g. "
    "one age group's share of the total poverty-count impact)."
)

set_seed(20260915)
programs = ["no_ss", "no_snap", "no_ctc", "no_housing"]
age_groups = ["Under 18", "18 to 64", "65+"]

# %%
logger.info(
    "One shared dataset, split by age group - 'Overall' is the StatCalculator "
    "over the whole thing, and each age group is the StatCalculator over its "
    "own filtered slice. Since 'sum' is additive over a partition of rows, "
    "Overall's total is guaranteed to equal the sum of the three age groups' "
    "own sums (unlike building each one from independent, unrelated data,  "
    "which would leave the 'total' label matching nothing on the chart)."
)
n_rows = 3_800
df = (
    RandomData(n_rows=n_rows, seed=generate_seed())
    .index("index")
    .float("impact_no_ss", -5, 0)
    .float("impact_no_snap", -3, 0)
    .float("impact_no_ctc", -2, 0)
    .float("impact_no_housing", -1, 0)
    .integer("age_group_id", 0, len(age_groups) - 1)
).to_df()

when = pl
for i, label in enumerate(age_groups):
    when = when.when(pl.col("age_group_id") == i).then(pl.lit(label))
df = df.with_columns(when.otherwise(pl.lit("Other")).alias("age_group")).drop(
    "age_group_id"
)
df = pl.concat(
    [
        df,
        bayes_bootstrap(
            n_rows=n_rows,
            n_draws=9,
            seed=generate_seed(),
            initial_weight_index=0,
            prefix="weight_",
        ),
    ],
    how="horizontal",
)

stats = Statistics(stats=["sum"], columns=[f"impact_{p}" for p in programs])
replicates = Replicates(weight_stub="weight_", n_replicates=8, bootstrap=True)

items = {
    label: StatCalculator(
        df.filter(pl.col("age_group") == label),
        statistics=stats,
        weight="weight_0",
        replicates=replicates,
    )
    for label in age_groups
}
items["Overall"] = StatCalculator(
    df, statistics=stats, weight="weight_0", replicates=replicates
)

for name, item in items.items():
    logger.info(name)
    item.print()

# %%
logger.info(
    "total_key excludes that entry from the stack and instead shows its own "
    "value as a text label next to each fully-stacked bar - here it lands "
    "right at the tip of each bar, since the age-group layers exactly sum "
    "to Overall by construction."
)
rename = {f"impact_{p}": p for p in programs}
fig_stacked_bar = plot.stacked_bar(
    items,
    column="sum",
    total_key="Overall",
    rename=rename,
    order=list(rename.values()),
    label_round_digits=0,
    x_axis_title="Change in number of people in poverty",
)
fig_stacked_bar.write_html(
    os.path.join(path_docs_figures, "stacked_bar.html"), include_plotlyjs="directory"
)

# %%
logger.info(
    "Layers aren't required to share a sign, and neither are whole "
    "categories - three cases side by side here: 'no_ss' keeps its "
    "negative (poverty-reducing) values for two age groups but flips "
    "positive for 'Under 18', so that one bar stacks in both directions at "
    "once (a block right of zero alongside blocks left of it); 'no_housing' "
    "is flipped positive for every age group, an all-plus bar entirely to "
    "the right; 'no_snap'/'no_ctc' stay all-negative, as before. Each "
    "total label lands on whichever side its own net sum ends up on."
)
df_mixed = df.with_columns(
    pl.when(pl.col("age_group") == "Under 18")
    .then(-pl.col("impact_no_ss"))
    .otherwise(pl.col("impact_no_ss"))
    .alias("impact_no_ss"),
    (-pl.col("impact_no_housing")).alias("impact_no_housing"),
)
items_mixed = {
    label: StatCalculator(
        df_mixed.filter(pl.col("age_group") == label),
        statistics=stats,
        weight="weight_0",
        replicates=replicates,
    )
    for label in age_groups
}
items_mixed["Overall"] = StatCalculator(
    df_mixed, statistics=stats, weight="weight_0", replicates=replicates
)
fig_stacked_bar_mixed = plot.stacked_bar(
    items_mixed,
    column="sum",
    total_key="Overall",
    rename=rename,
    order=list(rename.values()),
    label_round_digits=0,
    x_axis_title="Change in number of people in poverty",
)
fig_stacked_bar_mixed.write_html(
    os.path.join(path_docs_figures, "stacked_bar_mixed_sign.html"),
    include_plotlyjs="directory",
)

Two independent runs, each with and without a confidence band. Covers a 2-level tree (Run -> CI), a 3-level one (Run -> CI -> Year), labeling each level, and putting a level on its own row.

Figures: 2 levels · 3 levels

import os
import narwhals as nw
import polars as pl

from survey_kit.utilities.random import RandomData, set_seed, generate_seed
from survey_kit.statistics.calculator import StatCalculator
from survey_kit.statistics.statistics import Statistics
from survey_kit.statistics.replicates import Replicates
from survey_kit.statistics.bootstrap import bayes_bootstrap
from survey_kit import logger, config, plot

path_docs_plot = os.path.join(config.code_root, "..", "..", "docs", "tutorials", "plot")
path_docs_figures = os.path.join(path_docs_plot, "figures")
os.makedirs(path_docs_figures, exist_ok=True)

# %%
logger.info(
    "combine() puts several whole figures (from line()/quantiles()/coefplot()/ "
    "stacked_bar(), or any plotly Figure) into one HTML page, switched between "
    "by a dropdown per nesting level - as many levels deep as the dict you "
    "pass it. Picking a value at a deeper level is remembered when you switch "
    "a shallower one, and falls back to that branch's first option if the "
    "previous choice doesn't exist there. Each leaf figure's own internal "
    "group dropdown (see line_and_quantiles.py/coefplot.py) keeps working "
    "independently once it's shown - combine() only adds this outer layer."
)

set_seed(20260915)
n_rows = 1_500
n_replicates = 8


def make_quantiles_sc(seed: int) -> StatCalculator:
    df = (
        RandomData(n_rows=n_rows, seed=seed)
        .index("index")
        .integer("income", 0, 100_000)
        .integer("year", 2016, 2018)
    ).to_df()
    df = pl.concat(
        [
            df,
            bayes_bootstrap(
                n_rows=n_rows,
                n_draws=n_replicates + 1,
                seed=generate_seed(),
                initial_weight_index=0,
                prefix="weight_",
            ),
        ],
        how="horizontal",
    )
    stats = Statistics(stats=["q10", "q25", "q50", "q75", "q90"], columns=["income"])
    replicates = Replicates(
        weight_stub="weight_", n_replicates=n_replicates, bootstrap=True
    )
    return StatCalculator(
        df,
        statistics=stats,
        weight="weight_0",
        replicates=replicates,
        by={"year": ["year"]},
    )


# %%
logger.info(
    "Two independent StatCalculators (as if from two different runs/vintages) "
    "give us something worth switching between at the top level; each is "
    "plotted twice (with and without a confidence band), giving a 3-level "
    "tree: Run -> CI -> (nothing further, since quantiles() already puts every "
    "year on one figure)."
)
sc_run_a = make_quantiles_sc(generate_seed())
sc_run_b = make_quantiles_sc(generate_seed())

tree = {
    "Run A": {
        "With CI": plot.quantiles(sc_run_a, ci_level=0.95, ci_area=True),
        "No CI": plot.quantiles(sc_run_a),
    },
    "Run B": {
        "With CI": plot.quantiles(sc_run_b, ci_level=0.95, ci_area=True),
        "No CI": plot.quantiles(sc_run_b),
    },
}

# %%
logger.info(
    "Every leaf here already has its own (inert, single-group) dropdown from "
    "quantiles() - combine() adds the 'Run A'/'Run B' and 'With CI'/'No CI' "
    "switches on top, in one page. layout={'margin': {'t': 30}} trims Plotly's "
    "default top margin (~100px of otherwise-blank space above the plot, "
    "reserved for a title none of these figures use) - a stopgap until "
    "quantiles()/line() default to something less white-spacey on their own."
)
combined = plot.combine(tree, label=["Run:", "CI:"], layout={"margin": {"t": 30}})
combined.write_html(os.path.join(path_docs_figures, "combine.html"))

# %%
logger.info(
    "Nesting isn't limited to 2 levels - here's a 3-level tree (Run -> CI -> "
    "Year), built by giving quantiles() a `filter_expr` so each leaf covers "
    "just one year instead of all three. Switching 'Run' while you're looking "
    "at, say, 2017 keeps 2017 selected on the other run too, as long as 2017 "
    "exists there (it does here, since both runs cover the same years)."
)


def year_branches(
    sc: StatCalculator, ci_level: float | None, ci_area: bool = False
) -> dict:
    return {
        str(year): plot.quantiles(
            sc, ci_level=ci_level, ci_area=ci_area, filter_expr=nw.col("year") == year
        )
        for year in [2016, 2017, 2018]
    }


tree_deep = {
    "Run A": {
        "With CI": year_branches(sc_run_a, ci_level=0.95, ci_area=True),
        "No CI": year_branches(sc_run_a, ci_level=None),
    },
    "Run B": {
        "With CI": year_branches(sc_run_b, ci_level=0.95, ci_area=True),
        "No CI": year_branches(sc_run_b, ci_level=None),
    },
}

combined_deep = plot.combine(
    tree_deep,
    label=["Run:", "\nCI:", "\nYear:"],
    dropdowns_padding_left=50,
    layout={"margin": {"t": 30}},
)
combined_deep.write_html(os.path.join(path_docs_figures, "combine_3_levels.html"))
logger.info(os.path.join(path_docs_figures, "combine_3_levels.html"))

coefplot()/stacked_bar() have no group dropdown of their own, but their legend clicks are still shared by trace name - isolating a series/layer in one branch and switching to a sibling with the same name shows it isolated there too.

Figures: coefplot() · stacked_bar()

import os
import polars as pl

from survey_kit.utilities.random import RandomData, set_seed, generate_seed
from survey_kit.statistics.calculator import StatCalculator
from survey_kit.statistics.statistics import Statistics
from survey_kit.statistics.replicates import Replicates
from survey_kit.statistics.bootstrap import bayes_bootstrap
from survey_kit import logger, config, plot

path_docs_plot = os.path.join(config.code_root, "..", "..", "docs", "tutorials", "plot")
path_docs_figures = os.path.join(path_docs_plot, "figures")
os.makedirs(path_docs_figures, exist_ok=True)

# %%
logger.info(
    "coefplot() has no group dropdown of its own (unlike line()/quantiles() - "
    "see add_group_dropdown) and never calls add_group_dropdown, so its own "
    "legend-click toggle/isolate is plain, unwired Plotly - each figure's "
    "own legend state, not the shared name-keyed state that lets a "
    "line()/quantiles() sibling pick up the same isolate pattern when "
    "combine() switches to it. Two coefplots with the SAME series names "
    "(here: year) under two different combine() branches is the test - "
    "isolate '2018' in Region A, switch to Region B, switch back: does "
    "Region A still show only 2018, or did switching reset it?"
)

set_seed(20260915)
n_rows = 3_000
n_replicates = 8
subgroups = ["Male", "Female", "White", "Black"]
regions = ["Region A", "Region B"]


def make_coefplot(region: str, seed: int):
    df = (
        RandomData(n_rows=n_rows, seed=seed)
        .index("index")
        .integer("poverty", 0, 1)
        .integer("subgroup_id", 0, len(subgroups) - 1)
        .integer("year", 2016, 2018)
    ).to_df()

    when = pl
    for i, label in enumerate(subgroups):
        when = when.when(pl.col("subgroup_id") == i).then(pl.lit(label))
    df = df.with_columns(when.otherwise(pl.lit("Other")).alias("subgroup")).drop(
        "subgroup_id"
    )

    df = pl.concat(
        [
            df,
            bayes_bootstrap(
                n_rows=n_rows,
                n_draws=n_replicates + 1,
                seed=generate_seed(),
                initial_weight_index=0,
                prefix="weight_",
            ),
        ],
        how="horizontal",
    )

    stats = Statistics(stats=["mean"], columns=["poverty"])
    replicates = Replicates(
        weight_stub="weight_", n_replicates=n_replicates, bootstrap=True
    )
    sc = StatCalculator(
        df,
        statistics=stats,
        weight="weight_0",
        replicates=replicates,
        by={"group": ["subgroup", "year"]},
    )
    return plot.coefplot(
        sc,
        column="mean",
        ci_level=0.95,
        category="subgroup",
        series="year",
        x_axis_title=f"Poverty rate, {region}",
    )


tree = {region: make_coefplot(region, generate_seed()) for region in regions}

# %%
logger.info(
    "Both figures have the same trace names (2016/2017/2018 from series="
    "'year') - if you toggle/isolate one via the legend on Region A, switch "
    "to Region B, then back to Region A, check whether Region A's own "
    "toggle state survived (it should - it's the same DOM element, just "
    "hidden/shown - the open question is only whether Region B ALSO picked "
    "up the same isolate pattern when you first switched to it, the way a "
    "quantiles() sibling would)."
)
combined = plot.combine(tree, label="Region:", layout={"margin": {"t": 30}})
combined.write_html(os.path.join(path_docs_figures, "combine_two_coefplots.html"))
logger.info(os.path.join(path_docs_figures, "combine_two_coefplots.html"))
import os
import polars as pl

from survey_kit.utilities.random import RandomData, set_seed, generate_seed
from survey_kit.statistics.calculator import StatCalculator
from survey_kit.statistics.statistics import Statistics
from survey_kit.statistics.replicates import Replicates
from survey_kit.statistics.bootstrap import bayes_bootstrap
from survey_kit import logger, config, plot

path_docs_plot = os.path.join(config.code_root, "..", "..", "docs", "tutorials", "plot")
path_docs_figures = os.path.join(path_docs_plot, "figures")
os.makedirs(path_docs_figures, exist_ok=True)

# %%
logger.info(
    "Same question as combine_two_coefplots.py, for stacked_bar() instead - "
    "it also never calls add_group_dropdown, so its legend is plain, "
    "unwired Plotly too. Two stacked bars with the SAME layer names (age "
    "groups) under two different combine() branches: isolate/hide a layer "
    "via the legend on Scenario A, switch to Scenario B and back, and see "
    "whether Scenario A's own state survived and whether Scenario B picked "
    "up the same pattern when first switched to."
)

set_seed(20260915)
programs = ["no_ss", "no_snap", "no_ctc", "no_housing"]
age_groups = ["Under 18", "18 to 64", "65+"]
scenarios = ["Scenario A", "Scenario B"]


def make_stacked_bar(scenario: str, seed: int):
    n_rows = 3_800
    df = (
        RandomData(n_rows=n_rows, seed=seed)
        .index("index")
        .float("impact_no_ss", -5, 0)
        .float("impact_no_snap", -3, 0)
        .float("impact_no_ctc", -2, 0)
        .float("impact_no_housing", -1, 0)
        .integer("age_group_id", 0, len(age_groups) - 1)
    ).to_df()

    when = pl
    for i, label in enumerate(age_groups):
        when = when.when(pl.col("age_group_id") == i).then(pl.lit(label))
    df = df.with_columns(when.otherwise(pl.lit("Other")).alias("age_group")).drop(
        "age_group_id"
    )
    df = pl.concat(
        [
            df,
            bayes_bootstrap(
                n_rows=n_rows,
                n_draws=9,
                seed=generate_seed(),
                initial_weight_index=0,
                prefix="weight_",
            ),
        ],
        how="horizontal",
    )

    stats = Statistics(stats=["sum"], columns=[f"impact_{p}" for p in programs])
    replicates = Replicates(weight_stub="weight_", n_replicates=8, bootstrap=True)
    items = {
        label: StatCalculator(
            df.filter(pl.col("age_group") == label),
            statistics=stats,
            weight="weight_0",
            replicates=replicates,
        )
        for label in age_groups
    }
    items["Overall"] = StatCalculator(
        df, statistics=stats, weight="weight_0", replicates=replicates
    )
    rename = {f"impact_{p}": p for p in programs}
    return plot.stacked_bar(
        items,
        column="sum",
        total_key="Overall",
        rename=rename,
        order=list(rename.values()),
        label_round_digits=0,
        x_axis_title=f"Change in number of people in poverty, {scenario}",
    )


tree = {
    scenario: make_stacked_bar(scenario, generate_seed()) for scenario in scenarios
}

# %%
logger.info(
    "Both figures have the same layer/trace names (no_ss/no_snap/no_ctc/"
    "no_housing) - same check as the coefplot version: does a legend "
    "toggle on Scenario A survive switching away and back, and does "
    "switching to Scenario B pick up the same pattern the way a "
    "quantiles() sibling would?"
)
combined = plot.combine(tree, label="Scenario:", layout={"margin": {"t": 30}})
combined.write_html(os.path.join(path_docs_figures, "combine_two_stacked_bars.html"))
logger.info(os.path.join(path_docs_figures, "combine_two_stacked_bars.html"))