feat: take the arrests-over-time interval from the draws, not summed bounds
Deploy to git-pages / deploy (push) Successful in 16s
Deploy to git-pages / deploy (push) Successful in 16s
The band on the time-series chart summed each student group's own
count_lower/count_upper. That is not the interval of the district total: the sum
of eight 97.5th percentiles is the case where every group lands at its extreme
in the same draw, which is far less likely than any one group doing so. The
chart was overstating uncertainty by roughly a factor of two.
Now summed within each draw and summarized across draws — build_state_summary()'s
method from crdc-arrests/R/summarize_draws.R on a different axis, the same
operation poolBySex already performs. Clark County NV, 95% interval of the total:
wave draws-based summed bounds change
2015-16 215-276 (w 61) 172-322 (w 150) -60%
2017-18 175-230 (w 55) 134-269 (w 135) -60%
2021-22 116-164 (w 48) 78-188 (w 110) -57%
Medians barely move (245->246, 200->202, 137->138), which is the expected
signature: the sum of medians already approximates the median of the sum. It is
the tails that were wrong. Values verified against an offline DuckDB computation
over the same shards and then read back off the rendered SVG geometry.
Cost is very uneven, so it is not unconditional. Three waves is 0.31MB for
Nevada but 13.9MB for California and 14.3MB for Texas. Directory listings
(~1KB/year, and already fetched for part discovery) carry each part's byte size,
so useDistrictTotalDraws probes the total up front: under 3MB it fetches
automatically, above it the chart offers a button naming the real size
("Compute the exact interval from the draws (13.9 MB)"). Measured: Nevada adds
no perceptible time at all — the totals resolve in the same frame as the density
panel, since both wait on the duckdb-wasm engine that dominates a cold load —
and California's opt-in completes in 1.3s.
totalPerDraw refuses to sum unless every group has a complete draw set, and
fetchModelCounts now reports how many groups it dropped: for a per-group density
a missing group is a gap, but for a total it is an undercount that would shift
the whole interval down while looking entirely plausible.
The caption states which interval is on screen in either case, and names the
assumption behind the draws-based one — draw_id is renumbered per write batch
upstream, so summing at a draw index convolves what are effectively independent
draws. That is defensible for a posterior *predictive* total, whose observation
noise is independent across groups by construction and dominates the
parameter-level covariance, but it should be read as a predictive total rather
than a correlation-preserving contrast.
This commit is contained in:
@@ -11,6 +11,59 @@ import { MODEL_QUADRANT_LABEL } from '../utils/labels.js'
|
||||
* read as a pair.
|
||||
*/
|
||||
|
||||
/**
|
||||
* Says which interval is actually on screen, and offers the better one when it
|
||||
* hasn't been fetched. The distinction is not cosmetic: summing each group's
|
||||
* own bounds answers "what if every group hit its extreme at once", which is a
|
||||
* far wider claim than "how many arrests does the model think there were".
|
||||
*/
|
||||
function TotalIntervalNote({ totals, usingExact, onRequest }) {
|
||||
const base = { fontSize: '0.75rem', margin: 'var(--space-1) 0 0' }
|
||||
|
||||
if (usingExact) {
|
||||
return (
|
||||
<p style={{ ...base, color: 'var(--cv-ink-3)' }}>
|
||||
The modeled band is the 95% interval of the district <em>total</em>, taken from{' '}
|
||||
{(totals?.byYear && Object.values(totals.byYear).find(Boolean)?.nDraws?.toLocaleString()) || '500'}{' '}
|
||||
posterior predictive draws — summed across student groups within each draw, then
|
||||
summarized across draws.
|
||||
</p>
|
||||
)
|
||||
}
|
||||
|
||||
if (totals?.status === 'loading') {
|
||||
return <p style={{ ...base, color: 'var(--cv-ink-3)' }}>Computing the total from posterior draws…</p>
|
||||
}
|
||||
|
||||
const mb = totals?.bytes ? (totals.bytes / 1048576).toFixed(1) : null
|
||||
return (
|
||||
<p style={{ ...base, color: 'var(--cv-ink-3)' }}>
|
||||
The modeled band sums each student group’s own 95% bounds, which is wider than the
|
||||
interval of the total — it is the case where every group lands at its extreme in the same
|
||||
draw.{' '}
|
||||
{totals?.status === 'error' ? (
|
||||
<>The exact total could not be computed from the draws for this district.</>
|
||||
) : (
|
||||
onRequest && (
|
||||
<button type="button" onClick={onRequest} style={linkButton}>
|
||||
Compute the exact interval from the draws{mb ? ` (${mb} MB)` : ''}
|
||||
</button>
|
||||
)
|
||||
)}
|
||||
</p>
|
||||
)
|
||||
}
|
||||
|
||||
const linkButton = {
|
||||
background: 'none',
|
||||
border: 'none',
|
||||
padding: 0,
|
||||
font: 'inherit',
|
||||
color: 'var(--cv-navy-600)',
|
||||
textDecoration: 'underline',
|
||||
cursor: 'pointer',
|
||||
}
|
||||
|
||||
const WAVE_LABELS = { '15-16': '2015–16', '17-18': '2017–18', '21-22': '2021–22' }
|
||||
const HEADROOM = 26 // px reserved at the top of the plot so a point's rate label has room to sit above it
|
||||
|
||||
@@ -27,9 +80,21 @@ const MARGIN = { top: 20, right: 24, bottom: 52, left: 56 }
|
||||
// glyph is visually clipped and its label has nowhere to go.
|
||||
const X_PAD = 64
|
||||
|
||||
export default function ArrestsOverTime({ data, modelId }) {
|
||||
export default function ArrestsOverTime({ data, modelId, totals, onRequestExactTotals }) {
|
||||
// Prefer the interval computed from the draws (summed within each draw) over
|
||||
// the sum of each group's own bounds, which is not the interval of the total
|
||||
// and comes out systematically too wide.
|
||||
const exact = totals?.status === 'ready' ? totals.byYear : null
|
||||
const series = data.map((d) => {
|
||||
const iv = exact?.[d.year]
|
||||
return iv
|
||||
? { ...d, modeledMedian: iv.median, modeledLower: iv.lower, modeledUpper: iv.upper, exact: true }
|
||||
: { ...d, exact: false }
|
||||
})
|
||||
const usingExact = series.some((d) => d.exact)
|
||||
|
||||
const maxArrests = Math.max(
|
||||
...data.map((d) => Math.max(d.arrests, d.modeledUpper ?? 0)),
|
||||
...series.map((d) => Math.max(d.arrests, d.modeledUpper ?? 0)),
|
||||
1
|
||||
)
|
||||
const { ticks, niceMax } = niceTicks(maxArrests)
|
||||
@@ -43,7 +108,7 @@ export default function ArrestsOverTime({ data, modelId }) {
|
||||
|
||||
const plotLeft = margin.left + X_PAD
|
||||
const plotRight = margin.left + innerWidth - X_PAD
|
||||
const xScale = (i) => plotLeft + (i / Math.max(data.length - 1, 1)) * (plotRight - plotLeft)
|
||||
const xScale = (i) => plotLeft + (i / Math.max(series.length - 1, 1)) * (plotRight - plotLeft)
|
||||
const yScale = (val) => HEADROOM + (innerHeight - HEADROOM) * (1 - val / niceMax)
|
||||
|
||||
return (
|
||||
@@ -81,7 +146,7 @@ export default function ArrestsOverTime({ data, modelId }) {
|
||||
Total arrests
|
||||
</text>
|
||||
|
||||
{data.map((d, i) => (
|
||||
{series.map((d, i) => (
|
||||
<text key={d.year} x={xScale(i)} y={margin.top + yScale(0) + 18}
|
||||
textAnchor="middle" fontSize="0.62rem" fill="var(--cv-ink-2)">
|
||||
{WAVE_LABELS[d.year] || d.label}
|
||||
@@ -94,7 +159,7 @@ export default function ArrestsOverTime({ data, modelId }) {
|
||||
{/* Modeled point-range per wave — drawn first, directly under the
|
||||
observed marks it pairs with, and kept visually subordinate (thinner,
|
||||
no halo) so the observed series reads as the primary line. */}
|
||||
{data.map((d, i) => {
|
||||
{series.map((d, i) => {
|
||||
if (d.modeledMedian == null) return null
|
||||
const cx = xScale(i)
|
||||
const cyMedian = margin.top + yScale(d.modeledMedian)
|
||||
@@ -111,9 +176,9 @@ export default function ArrestsOverTime({ data, modelId }) {
|
||||
})}
|
||||
|
||||
{/* Observed line */}
|
||||
{data.length > 1 && (
|
||||
{series.length > 1 && (
|
||||
<polyline
|
||||
points={data.map((d, i) => `${xScale(i)},${margin.top + yScale(d.arrests)}`).join(' ')}
|
||||
points={series.map((d, i) => `${xScale(i)},${margin.top + yScale(d.arrests)}`).join(' ')}
|
||||
fill="none" stroke={OBSERVED_MARK_COLOR} strokeWidth={2}
|
||||
strokeLinejoin="round" strokeLinecap="round"
|
||||
/>
|
||||
@@ -122,7 +187,7 @@ export default function ArrestsOverTime({ data, modelId }) {
|
||||
{/* Observed diamonds + rate-per-1k labels. X_PAD keeps the end markers
|
||||
clear of the plot edges, so every label can be centred over its own
|
||||
point instead of being anchored outward to avoid an overflow. */}
|
||||
{data.map((d, i) => {
|
||||
{series.map((d, i) => {
|
||||
const cx = xScale(i)
|
||||
const cy = margin.top + yScale(d.arrests)
|
||||
const ratePerK = d.enroll > 0 ? (d.arrests / (d.enroll / 1000)).toFixed(2) : '0.00'
|
||||
@@ -145,15 +210,17 @@ export default function ArrestsOverTime({ data, modelId }) {
|
||||
|
||||
<ChartLegend items={[
|
||||
{ shape: 'diamond', color: OBSERVED_MARK_COLOR, label: 'Observed' },
|
||||
// 95%, not 90%: these bars are the API's count_lower/count_upper, and
|
||||
// validate_interval() defaults to 95 (the app never passes interval=).
|
||||
// 95%, not 90%: the API's count_lower/count_upper are a 95% interval
|
||||
// (validate_interval() defaults to 95), and the draws-based total is
|
||||
// computed at the same mass, so the two are directly comparable.
|
||||
{ shape: 'dot', color: MODELED_AGGREGATE_COLOR, label: 'Modeled (median + 95% interval)' },
|
||||
]} />
|
||||
|
||||
<TotalIntervalNote totals={totals} usingExact={usingExact} onRequest={onRequestExactTotals} />
|
||||
|
||||
<p style={{ fontSize: '0.75rem', color: 'var(--cv-ink-3)', marginTop: 'var(--space-1)' }}>
|
||||
Rate per 1,000 students labeled above each observed point. The modeled total sums each
|
||||
selected group's median/interval independently, which approximates but is not exactly the
|
||||
median of the combined total. Data from CRDC waves 2015–16 through 2021–22.
|
||||
Rate per 1,000 students labeled above each observed point. Data from CRDC waves 2015–16
|
||||
through 2021–22.
|
||||
</p>
|
||||
</div>
|
||||
)
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import { useCallback, useEffect, useMemo, useState } from 'react'
|
||||
import * as api from '../hooks/useApi.js'
|
||||
import { MODEL_QUADRANTS } from '../hooks/useApi.js'
|
||||
import { useDrawDistribution } from '../hooks/useDrawDistribution.js'
|
||||
import { useDistrictTotalDraws, useDrawDistribution } from '../hooks/useDrawDistribution.js'
|
||||
import {
|
||||
buildDisplayGroups,
|
||||
defaultDiffPair,
|
||||
@@ -20,6 +20,11 @@ const DEFAULT_SPEC = 'unified_m4_mod' // three-year + referral rate
|
||||
const CURRENT_WAVE = '21-22'
|
||||
const QUADRANT_MODELS = MODEL_QUADRANTS.map((q) => q.model)
|
||||
|
||||
// Below this, the draws for all three waves are fetched without asking. Sized
|
||||
// from the real shards: Nevada's three waves are 0.31MB, California's 13.9MB
|
||||
// and Texas's 14.3MB, so 3MB cleanly separates "free" from "worth a click".
|
||||
const AUTO_TOTAL_DRAW_BYTES = 3 * 1024 * 1024
|
||||
|
||||
/**
|
||||
* The results page: what was observed (summary table), what the model says the
|
||||
* rate is (density panel), and what it says about the gap between two groups
|
||||
@@ -43,6 +48,7 @@ export default function ChartPanel({ district, state }) {
|
||||
// made in one mode must not be reapplied in the other.
|
||||
const [selection, setSelection] = useState(null)
|
||||
const [pairSelection, setPairSelection] = useState(null)
|
||||
const [exactTotals, setExactTotals] = useState(false)
|
||||
|
||||
// ——— Fetch: the three waves for the time series ———
|
||||
useEffect(() => {
|
||||
@@ -54,6 +60,7 @@ export default function ChartPanel({ district, state }) {
|
||||
setSelection(null)
|
||||
setPairSelection(null)
|
||||
setCompareAll(false)
|
||||
setExactTotals(false)
|
||||
|
||||
async function run() {
|
||||
const waves = {}
|
||||
@@ -154,6 +161,22 @@ export default function ChartPanel({ district, state }) {
|
||||
year: CURRENT_WAVE,
|
||||
})
|
||||
|
||||
// The time series' band, computed from the draws rather than by summing each
|
||||
// group's own bounds. Fetched automatically only where it is cheap — three
|
||||
// waves is 0.3MB for Nevada but 13.9MB for California — otherwise the chart
|
||||
// offers it as an explicit choice with the size shown.
|
||||
const totals = useDistrictTotalDraws({
|
||||
leaid: district.leaid,
|
||||
state,
|
||||
model: WAVE_MODEL,
|
||||
years: ALL_WAVES,
|
||||
enabled: exactTotals,
|
||||
})
|
||||
|
||||
useEffect(() => {
|
||||
if (totals.bytes > 0 && totals.bytes <= AUTO_TOTAL_DRAW_BYTES) setExactTotals(true)
|
||||
}, [totals.bytes])
|
||||
|
||||
if (loading || !waveData) return <LoadingCharts />
|
||||
|
||||
const timeSeriesData = ALL_WAVES.map((year) => {
|
||||
@@ -223,7 +246,12 @@ export default function ChartPanel({ district, state }) {
|
||||
onPairChange={handlePairChange}
|
||||
/>
|
||||
|
||||
<ArrestsOverTime data={timeSeriesData} modelId={WAVE_MODEL} />
|
||||
<ArrestsOverTime
|
||||
data={timeSeriesData}
|
||||
modelId={WAVE_MODEL}
|
||||
totals={totals}
|
||||
onRequestExactTotals={() => setExactTotals(true)}
|
||||
/>
|
||||
</div>
|
||||
|
||||
<div
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import { useEffect, useMemo, useState } from 'react'
|
||||
import { getDb } from '../utils/duckdbClient.js'
|
||||
import { groupKey, isCompleteDrawSet } from '../utils/drawGroups.js'
|
||||
import { totalInterval, totalPerDraw } from '../utils/districtTotal.js'
|
||||
|
||||
const HF_DATASET = 'civilytics/crdc-school-arrest-rates'
|
||||
const HF_BASE = `https://huggingface.co/datasets/${HF_DATASET}/resolve/main/parquet`
|
||||
@@ -24,6 +25,57 @@ function shardDir(model, year, state) {
|
||||
return `model_id=${model}/YEAR=${year}/LEA_STATE=${state}`
|
||||
}
|
||||
|
||||
// Directory listings are ~1KB and carry each part's byte size, so the UI can
|
||||
// tell the reader what a fetch will cost before committing to it. Cached
|
||||
// separately from the buffers: listing a shard is cheap, downloading it is not.
|
||||
const listingCache = new Map()
|
||||
|
||||
/**
|
||||
* @returns {Promise<Array<{part: number, size: number}>>} parquet parts, in order
|
||||
*/
|
||||
function listShardEntries(model, year, state) {
|
||||
const dir = shardDir(model, year, state)
|
||||
if (!listingCache.has(dir)) {
|
||||
listingCache.set(
|
||||
dir,
|
||||
(async () => {
|
||||
try {
|
||||
const res = await fetch(`${HF_TREE}/${dir}`)
|
||||
if (!res.ok) throw new Error(`tree listing HTTP ${res.status}`)
|
||||
const entries = await res.json()
|
||||
return entries
|
||||
.filter((e) => e?.type === 'file' && /\/data_\d+\.parquet$/.test(e.path || ''))
|
||||
.map((e) => ({ part: Number(e.path.match(/data_(\d+)\.parquet$/)[1]), size: e.size || 0 }))
|
||||
.sort((a, b) => a.part - b.part)
|
||||
} catch (err) {
|
||||
listingCache.delete(dir)
|
||||
throw err
|
||||
}
|
||||
})(),
|
||||
)
|
||||
}
|
||||
return listingCache.get(dir)
|
||||
}
|
||||
|
||||
/**
|
||||
* Total bytes of the draw shards for one model across several years — what a
|
||||
* "compute this from the draws" action will actually download.
|
||||
*
|
||||
* @returns {Promise<number>} bytes, or 0 if the size can't be determined
|
||||
*/
|
||||
export async function measureShardBytes(model, years, state) {
|
||||
try {
|
||||
const sizes = await Promise.all(
|
||||
years.map(async (year) =>
|
||||
(await listShardEntries(model, year, state)).reduce((sum, p) => sum + p.size, 0),
|
||||
),
|
||||
)
|
||||
return sizes.reduce((sum, n) => sum + n, 0)
|
||||
} catch {
|
||||
return 0
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Lists the parquet parts published for one (model, year, state).
|
||||
*
|
||||
@@ -42,13 +94,7 @@ function shardDir(model, year, state) {
|
||||
async function listShardParts(model, year, state) {
|
||||
const dir = shardDir(model, year, state)
|
||||
try {
|
||||
const res = await fetch(`${HF_TREE}/${dir}`)
|
||||
if (!res.ok) throw new Error(`tree listing HTTP ${res.status}`)
|
||||
const entries = await res.json()
|
||||
const parts = entries
|
||||
.filter((e) => e?.type === 'file' && /\/data_\d+\.parquet$/.test(e.path || ''))
|
||||
.map((e) => ({ part: Number(e.path.match(/data_(\d+)\.parquet$/)[1]), path: e.path }))
|
||||
.sort((a, b) => a.part - b.part)
|
||||
const parts = await listShardEntries(model, year, state)
|
||||
if (parts.length > 0) return parts.map((p) => p.part)
|
||||
throw new Error('tree listing contained no parquet parts')
|
||||
} catch (err) {
|
||||
@@ -139,12 +185,18 @@ async function fetchModelCounts(db, { leaid, state, model, year }) {
|
||||
// a short draw set: its density and interval would be computed off a
|
||||
// biased subsample and look identical on screen to a complete one.
|
||||
const counts = {}
|
||||
let dropped = 0
|
||||
for (const [key, arr] of Object.entries(raw)) {
|
||||
if (isCompleteDrawSet(arr, nDraws)) counts[key] = arr
|
||||
else console.warn('useDrawDistribution: dropping incomplete draw set for', { leaid, model, key, got: arr.length, want: nDraws })
|
||||
else {
|
||||
dropped += 1
|
||||
console.warn('useDrawDistribution: dropping incomplete draw set for', { leaid, model, key, got: arr.length, want: nDraws })
|
||||
}
|
||||
}
|
||||
|
||||
return Object.keys(counts).length > 0 ? { counts, nDraws } : null
|
||||
// `dropped` matters to any consumer that aggregates *across* groups (the
|
||||
// district total): a missing group there is an undercount, not a gap.
|
||||
return Object.keys(counts).length > 0 ? { counts, nDraws, dropped } : null
|
||||
} finally {
|
||||
if (conn) await conn.close()
|
||||
}
|
||||
@@ -257,3 +309,89 @@ export function useDrawDistribution({ leaid, state, models, year }) {
|
||||
|
||||
return { status, byModel, nDraws }
|
||||
}
|
||||
|
||||
/**
|
||||
* The district's total arrest count per wave, with a 95% interval computed from
|
||||
* the draws — summed within each draw, then summarized across draws.
|
||||
*
|
||||
* Gated behind `enabled` because the cost is wildly uneven: three waves of
|
||||
* Nevada is 0.3MB, but California is 13.9MB and Texas 14.3MB. `bytes` is
|
||||
* probed from the directory listings up front (~1KB per year) so the caller can
|
||||
* either fetch automatically when it's cheap or show the reader the price
|
||||
* first.
|
||||
*
|
||||
* @param {{leaid: string, state: string, model: string, years: string[], enabled: boolean}} params
|
||||
* @returns {{status: 'idle'|'loading'|'ready'|'error',
|
||||
* byYear: Record<string, {lower:number, median:number, upper:number, nDraws:number}> | null,
|
||||
* bytes: number}}
|
||||
*/
|
||||
export function useDistrictTotalDraws({ leaid, state, model, years, enabled }) {
|
||||
const [status, setStatus] = useState('idle')
|
||||
const [byYear, setByYear] = useState(null)
|
||||
const [bytes, setBytes] = useState(0)
|
||||
|
||||
const yearsSignature = (years || []).join(',')
|
||||
|
||||
// Probe sizes regardless of `enabled` — this is what lets the UI decide.
|
||||
useEffect(() => {
|
||||
let cancelled = false
|
||||
setBytes(0)
|
||||
if (!state || !model || !yearsSignature) return
|
||||
measureShardBytes(model, yearsSignature.split(','), state).then((n) => {
|
||||
if (!cancelled) setBytes(n)
|
||||
})
|
||||
return () => {
|
||||
cancelled = true
|
||||
}
|
||||
}, [state, model, yearsSignature])
|
||||
|
||||
useEffect(() => {
|
||||
let cancelled = false
|
||||
setStatus(enabled ? 'loading' : 'idle')
|
||||
setByYear(null)
|
||||
|
||||
if (!enabled || !leaid || !state || !model || !yearsSignature) return
|
||||
|
||||
async function run() {
|
||||
try {
|
||||
const db = await getDb()
|
||||
const yearList = yearsSignature.split(',')
|
||||
const results = await Promise.all(
|
||||
yearList.map(async (year) => {
|
||||
try {
|
||||
const model_ = await fetchModelCounts(db, { leaid, state, model, year })
|
||||
if (!model_ || model_.dropped > 0) return [year, null]
|
||||
const interval = totalInterval(totalPerDraw(model_.counts, model_.nDraws))
|
||||
return [year, interval]
|
||||
} catch (err) {
|
||||
console.error(`useDistrictTotalDraws: ${year} failed:`, err)
|
||||
return [year, null]
|
||||
}
|
||||
}),
|
||||
)
|
||||
if (cancelled) return
|
||||
const next = Object.fromEntries(results)
|
||||
if (Object.values(next).some((v) => v !== null)) {
|
||||
setByYear(next)
|
||||
setStatus('ready')
|
||||
} else {
|
||||
setByYear(null)
|
||||
setStatus('error')
|
||||
}
|
||||
} catch (err) {
|
||||
console.error('useDistrictTotalDraws failed:', err)
|
||||
if (!cancelled) {
|
||||
setByYear(null)
|
||||
setStatus('error')
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
run()
|
||||
return () => {
|
||||
cancelled = true
|
||||
}
|
||||
}, [leaid, state, model, yearsSignature, enabled])
|
||||
|
||||
return { status, byYear, bytes }
|
||||
}
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
/**
|
||||
* The district-wide arrest total, taken from the posterior predictive draws.
|
||||
*
|
||||
* Why this exists: the time-series chart used to draw its band by summing each
|
||||
* student group's own `count_lower`/`count_upper`. That is not the interval of
|
||||
* the total. The sum of per-group 97.5th percentiles is the value you would see
|
||||
* if *every* group simultaneously landed at its own extreme in the same draw,
|
||||
* which is far less likely than any one group doing so — so the band came out
|
||||
* systematically too wide. Summing within each draw and taking quantiles of the
|
||||
* resulting totals answers the actual question: how many arrests does the model
|
||||
* think this district had?
|
||||
*
|
||||
* This is `build_state_summary()`'s method from crdc-arrests/R/summarize_draws.R
|
||||
* applied on a different axis — sum inside the draw, summarize across draws —
|
||||
* the same operation `poolBySex` performs for sex pooling.
|
||||
*
|
||||
* One caveat to carry into any caption. Upstream, `draw_id` is renumbered per
|
||||
* write batch, so draw *k* of one group is not the same posterior sample as
|
||||
* draw *k* of another; measured cross-group correlation is ≈ 0.02. Summing at a
|
||||
* draw index therefore convolves what are effectively independent draws. That
|
||||
* is a reasonable model here — these are *posterior predictive* draws, whose
|
||||
* observation noise is independent across groups by construction and dominates
|
||||
* the parameter-level covariance — but it does mean the result should be read
|
||||
* as a predictive total, not as a contrast that preserves parameter
|
||||
* correlation. It is still much closer to the truth than summing bounds.
|
||||
*/
|
||||
|
||||
import { quantile } from './kde.js'
|
||||
import { isCompleteDrawSet } from './drawGroups.js'
|
||||
|
||||
/** Matches the API's default interval, so the two are directly comparable. */
|
||||
export const TOTAL_INTERVAL_MASS = 0.95
|
||||
|
||||
/**
|
||||
* District total at each draw index: the sum across every student group.
|
||||
*
|
||||
* Returns null unless *every* group has a complete draw set. A group silently
|
||||
* missing from the sum would undercount the total at every draw and shift the
|
||||
* whole interval down, which is indistinguishable on screen from a real result.
|
||||
*
|
||||
* @param {Record<string, number[]> | null | undefined} countsByGroup
|
||||
* @param {number} nDraws
|
||||
* @returns {number[] | null}
|
||||
*/
|
||||
export function totalPerDraw(countsByGroup, nDraws) {
|
||||
const groups = Object.values(countsByGroup || {})
|
||||
if (!groups.length || !(nDraws > 0)) return null
|
||||
if (!groups.every((counts) => isCompleteDrawSet(counts, nDraws))) return null
|
||||
|
||||
const totals = new Array(nDraws).fill(0)
|
||||
for (const counts of groups) {
|
||||
for (let i = 0; i < nDraws; i++) totals[i] += counts[i]
|
||||
}
|
||||
return totals
|
||||
}
|
||||
|
||||
/**
|
||||
* @param {number[] | null | undefined} totals - district total per draw
|
||||
* @returns {{lower: number, median: number, upper: number, nDraws: number} | null}
|
||||
*/
|
||||
export function totalInterval(totals) {
|
||||
if (!totals?.length) return null
|
||||
const tail = (1 - TOTAL_INTERVAL_MASS) / 2
|
||||
return {
|
||||
lower: quantile(totals, tail),
|
||||
median: quantile(totals, 0.5),
|
||||
upper: quantile(totals, 1 - tail),
|
||||
nDraws: totals.length,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,86 @@
|
||||
import { test } from 'node:test'
|
||||
import assert from 'node:assert/strict'
|
||||
import { TOTAL_INTERVAL_MASS, totalInterval, totalPerDraw } from './districtTotal.js'
|
||||
|
||||
// ——— totalPerDraw ———
|
||||
|
||||
test('totalPerDraw: sums every group at the same draw index', () => {
|
||||
const counts = { BL_F: [1, 2, 3], BL_M: [10, 20, 30], WH_F: [100, 200, 300] }
|
||||
assert.deepEqual(totalPerDraw(counts, 3), [111, 222, 333])
|
||||
})
|
||||
|
||||
test('totalPerDraw: a single group is its own total', () => {
|
||||
assert.deepEqual(totalPerDraw({ BL_M: [4, 5] }, 2), [4, 5])
|
||||
})
|
||||
|
||||
test('totalPerDraw: null when any group is short of nDraws', () => {
|
||||
// Summing a 2-draw group into a 3-draw total would silently undercount the
|
||||
// last draw, biasing the whole interval downward.
|
||||
assert.equal(totalPerDraw({ BL_F: [1, 2, 3], BL_M: [1, 2] }, 3), null)
|
||||
})
|
||||
|
||||
test('totalPerDraw: null when a group has a hole', () => {
|
||||
const holey = [1, 2, 3]
|
||||
delete holey[1]
|
||||
assert.equal(totalPerDraw({ BL_F: holey }, 3), null)
|
||||
})
|
||||
|
||||
test('totalPerDraw: null for empty or missing input', () => {
|
||||
assert.equal(totalPerDraw({}, 500), null)
|
||||
assert.equal(totalPerDraw(null, 500), null)
|
||||
assert.equal(totalPerDraw({ BL_F: [1] }, 0), null)
|
||||
})
|
||||
|
||||
test('totalPerDraw: does not mutate its input', () => {
|
||||
const counts = { BL_F: [1, 2], BL_M: [3, 4] }
|
||||
totalPerDraw(counts, 2)
|
||||
assert.deepEqual(counts, { BL_F: [1, 2], BL_M: [3, 4] })
|
||||
})
|
||||
|
||||
// ——— totalInterval ———
|
||||
|
||||
test('totalInterval: median and 95% bounds from the draw totals', () => {
|
||||
const totals = Array.from({ length: 1001 }, (_, i) => i) // 0…1000
|
||||
const iv = totalInterval(totals)
|
||||
// Linear-interpolated quantiles, so compare with a tolerance rather than
|
||||
// exactly: p=0.025 over 0…1000 lands on 25 ± float noise.
|
||||
const near = (a, b) => Math.abs(a - b) < 1e-9
|
||||
assert.ok(near(iv.median, 500), `median ${iv.median}`)
|
||||
assert.ok(near(iv.lower, 25), `lower ${iv.lower}`)
|
||||
assert.ok(near(iv.upper, 975), `upper ${iv.upper}`)
|
||||
assert.equal(iv.nDraws, 1001)
|
||||
})
|
||||
|
||||
test('totalInterval: uses a 95% mass, matching the API convention', () => {
|
||||
assert.equal(TOTAL_INTERVAL_MASS, 0.95)
|
||||
})
|
||||
|
||||
test('totalInterval: lower <= median <= upper', () => {
|
||||
const totals = Array.from({ length: 500 }, (_, i) => (i * 7919) % 331)
|
||||
const iv = totalInterval(totals)
|
||||
assert.ok(iv.lower <= iv.median && iv.median <= iv.upper)
|
||||
})
|
||||
|
||||
test('totalInterval: a degenerate total collapses to a point', () => {
|
||||
const iv = totalInterval(new Array(500).fill(12))
|
||||
assert.deepEqual([iv.lower, iv.median, iv.upper], [12, 12, 12])
|
||||
})
|
||||
|
||||
test('totalInterval: null for empty or missing totals', () => {
|
||||
assert.equal(totalInterval([]), null)
|
||||
assert.equal(totalInterval(null), null)
|
||||
})
|
||||
|
||||
test('totalInterval: is narrower than summing the groups own bounds', () => {
|
||||
// The property that motivates this whole path. Two independent groups, each
|
||||
// roughly uniform on 0..100: summing each group's 97.5th percentile gives
|
||||
// ~200, but the 97.5th percentile of the *sum* is well below that, because
|
||||
// both groups landing at their extreme in the same draw is rare.
|
||||
const n = 4000
|
||||
const a = Array.from({ length: n }, (_, i) => (i * 37) % 101)
|
||||
const b = Array.from({ length: n }, (_, i) => (i * 61) % 101)
|
||||
const iv = totalInterval(totalPerDraw({ A: a, B: b }, n))
|
||||
const summedBounds = { lower: 0, upper: 100 + 100 }
|
||||
assert.ok(iv.upper < summedBounds.upper, `sum-of-quantiles ${summedBounds.upper} vs quantile-of-sum ${iv.upper}`)
|
||||
assert.ok(iv.upper - iv.lower < summedBounds.upper - summedBounds.lower)
|
||||
})
|
||||
Reference in New Issue
Block a user