Fix ridgeline chart to render smooth density curves instead of individual columns
Deploy to git-pages / deploy (push) Successful in 13s
Deploy to git-pages / deploy (push) Successful in 13s
- Rewrite RateDensityRidgeline.jsx with proper histogram-to-density conversion - Build top-edge points from binned draws, then smooth with d3.curveBasis - Construct complete area path: bottom edge + smoothed top + close - Position observed rate diamonds above the ridge peak
This commit is contained in:
+194
-200
@@ -3,83 +3,67 @@ import * as d3 from 'd3'
|
||||
|
||||
const MODEL_LABELS = {
|
||||
unified_m1_mod: 'One-year, baseline',
|
||||
unified_m2_mod: 'One-year + covariate',
|
||||
unified_m2_mod: 'One-year + covariate',
|
||||
unified_m3_mod: 'Three-year, baseline',
|
||||
unified_m4_mod: 'Three-year + covariate',
|
||||
}
|
||||
|
||||
const RACE_LABELS = {
|
||||
WH: 'White',
|
||||
BL: 'Black',
|
||||
BL: 'Black',
|
||||
HI: 'Hispanic',
|
||||
AM: 'American Indian/\nAlaska Native'
|
||||
AM: 'American Indian / Alaska Native'
|
||||
}
|
||||
|
||||
export default function RateDensityRidgeline({ quadData, rateByGroup }) {
|
||||
const [selectedModel, setSelectedModel] = useState('unified_m2_mod')
|
||||
const svgRef = useRef(null)
|
||||
|
||||
// Colors matching Civilytics palette
|
||||
const colors = {
|
||||
navy: '#000a9b',
|
||||
teal: '#0791b6',
|
||||
danger: '#c92d0e',
|
||||
ink: '#222222'
|
||||
}
|
||||
|
||||
useEffect(() => {
|
||||
if (!quadData || !selectedModel || !quadData[selectedModel]) return
|
||||
|
||||
const rows = quadData[selectedModel] || []
|
||||
|
||||
// Get unique groups for this model
|
||||
const groups = [...new Map(rows.map(r => [`${r.race}-${r.sex}`, {
|
||||
race: r.race,
|
||||
sex: r.sex,
|
||||
label: `${RACE_LABELS[r.race] || r.race} ${r.sex === 'F' ? 'Female' : 'Male'}`
|
||||
}])).values()]
|
||||
|
||||
// Build data for ridgeline
|
||||
const plotData = groups.map(g => {
|
||||
const row = rows.find(r => r.race === g.race && r.sex === g.sex)
|
||||
// Get unique groups for this model (race × sex)
|
||||
const groupKeys = [...new Set(rows.map(r => `${r.race}-${r.sex}`))]
|
||||
|
||||
// Build data: for each group, estimate density from the interval
|
||||
const plotData = groupKeys.map(key => {
|
||||
const row = rows.find(r => `${r.race}-${r.sex}` === key)
|
||||
if (!row) return null
|
||||
|
||||
// Generate histogram from draws or use rate bounds to approximate
|
||||
const enroll = row.stu_enroll || 1
|
||||
|
||||
// Create density-like data from the interval
|
||||
const rateMedian = (row.rate_median || 0) * 1000
|
||||
const rateLower = (row.rate_lower || 0) * 1000
|
||||
const rateUpper = (row.rate_upper || 0) * 1000
|
||||
|
||||
// Generate synthetic histogram from the interval (normal approximation)
|
||||
const sd = (rateUpper - rateLower) / 3.29 // 90% interval
|
||||
const draws = Array.from({ length: 100 }, () =>
|
||||
Math.max(0, d3.randomNormal(rateMedian, sd || 0.1)())
|
||||
)
|
||||
|
||||
// Observed rate
|
||||
const observedRate = (row.observed_arrests || 0) / ((enroll || 1) / 1000)
|
||||
|
||||
return {
|
||||
...g,
|
||||
draws: draws.sort(d3.ascending),
|
||||
observedRate,
|
||||
rateMedian
|
||||
}
|
||||
}).filter(Boolean)
|
||||
|
||||
// Observed rates from rateByGroup for comparison (single model only)
|
||||
const observedData = groups.map(g => {
|
||||
const groupData = rateByGroup.find(r => r.race === g.race && r.sex === g.sex)
|
||||
return groupData ? { ...g, observedRate: groupData.observedRate } : null
|
||||
const race = row.race || 'WH'
|
||||
const sex = row.sex || 'F'
|
||||
const label = `${RACE_LABELS[race] || race} ${sex === 'F' ? 'Female' : 'Male'}`
|
||||
|
||||
// Convert rate estimates to per-1000 scale
|
||||
const enroll = row.stu_enroll || 1
|
||||
const obsArrests = row.observed_arrests || 0
|
||||
const observedRate = enroll > 0 ? (obsArrests / enroll) * 1000 : 0
|
||||
|
||||
// Model rate estimates in per-1k units
|
||||
const rateMedian = ((row.rate_median || 0) * 1000) / (enroll > 0 ? 1 : 1)
|
||||
const rateLower = ((row.rate_lower || 0) * 1000) / (enroll > 0 ? 1 : 1)
|
||||
const rateUpper = ((row.rate_upper || 0) * 1000) / (enroll > 0 ? 1 : 1)
|
||||
|
||||
// Generate draws from normal approximation of the posterior interval
|
||||
// Use 90% CI: sd = (upper - lower) / 3.29
|
||||
const intervalWidth = Math.abs(rateUpper - rateLower) || 0.1
|
||||
const sd = intervalWidth / 3.29
|
||||
|
||||
// Generate a larger sample for smoother density estimation
|
||||
const draws = Array.from({ length: 500 }, () =>
|
||||
Math.max(0, d3.randomNormal(rateMedian, sd)() )
|
||||
).sort(d3.ascending)
|
||||
|
||||
return { key, label, race, sex, rateMedian, observedRate, draws }
|
||||
}).filter(Boolean)
|
||||
|
||||
// Set up dimensions
|
||||
const container = svgRef.current.parentElement
|
||||
const container = svgRef.current?.parentElement
|
||||
const width = Math.min(500, container?.clientWidth || 500)
|
||||
const height = Math.max(300, groups.length * 50 + 80)
|
||||
const margin = { top: 40, right: 30, bottom: 60, left: 80 }
|
||||
const height = Math.max(320, plotData.length * 60 + 80)
|
||||
const margin = { top: 40, right: 30, bottom: 70, left: 90 }
|
||||
const innerWidth = width - margin.left - margin.right
|
||||
const innerHeight = height - margin.top - margin.bottom
|
||||
|
||||
@@ -91,41 +75,34 @@ export default function RateDensityRidgeline({ quadData, rateByGroup }) {
|
||||
.attr('height', height)
|
||||
.attr('viewBox', `0 0 ${width} ${height}`)
|
||||
|
||||
// Create scales
|
||||
const allRates = plotData.flatMap(d => d.draws).concat(observedData.map(d => d.observedRate))
|
||||
const xMax = Math.min(Math.max(...allRates, 0), 20) || 5 // Cap at 20 per 1k for readability
|
||||
|
||||
// Compute x-scale domain from all data (draws + observed rates)
|
||||
const allDrawValues = plotData.flatMap(d => d.draws).filter(v => v <= 25)
|
||||
const maxRate = Math.min(Math.max(...allDrawValues, ...plotData.map(d => d.observedRate), 0.1), 25)
|
||||
|
||||
const xScale = d3.scaleLinear()
|
||||
.domain([0, xMax])
|
||||
.domain([0, maxRate])
|
||||
.range([margin.left, margin.left + innerWidth])
|
||||
|
||||
// Y-scale: one band per group (race × sex)
|
||||
const yScale = d3.scaleBand()
|
||||
.domain(plotData.map(d => d.label))
|
||||
.range([margin.top, margin.top + innerHeight])
|
||||
.padding(0.3)
|
||||
.padding(0.25)
|
||||
|
||||
// X-axis
|
||||
svg.append('g')
|
||||
.attr('transform', `translate(0,${margin.top + innerHeight})`)
|
||||
.call(d3.axisBottom(xScale).ticks(10))
|
||||
.call(g => g.select('.domain').attr('stroke', '#ddd'))
|
||||
.call(g => g.selectAll('.tick line').attr('stroke', '#eee'))
|
||||
.call(g => g.selectAll('.tick text').attr('fill', '#666'))
|
||||
|
||||
// X-axis label
|
||||
// Title
|
||||
svg.append('text')
|
||||
.attr('x', margin.left + innerWidth / 2)
|
||||
.attr('y', height - 10)
|
||||
.attr('y', 18)
|
||||
.attr('text-anchor', 'middle')
|
||||
.attr('font-size', '0.75rem')
|
||||
.attr('font-size', '0.7rem')
|
||||
.attr('fill', '#666')
|
||||
.text('Predicted arrests per 1,000 students')
|
||||
.text(`Posterior predicted arrests per 1,000 students — ${MODEL_LABELS[selectedModel]}`)
|
||||
|
||||
// Grid lines
|
||||
svg.append('g')
|
||||
.attr('class', 'grid')
|
||||
.selectAll('line')
|
||||
.data(xScale.ticks(10))
|
||||
.data(xScale.ticks(8))
|
||||
.join('line')
|
||||
.attr('x1', d => xScale(d))
|
||||
.attr('x2', d => xScale(d))
|
||||
@@ -134,140 +111,158 @@ export default function RateDensityRidgeline({ quadData, rateByGroup }) {
|
||||
.attr('stroke', '#f0f0f0')
|
||||
.attr('stroke-width', 1)
|
||||
|
||||
// Generate density polygons for each group
|
||||
plotData.forEach((d, i) => {
|
||||
const yPos = yScale(d.label)
|
||||
const rowHeight = yScale.bandwidth()
|
||||
|
||||
if (!yPos) return
|
||||
// X-axis
|
||||
svg.append('g')
|
||||
.attr('transform', `translate(0,${margin.top + innerHeight})`)
|
||||
.call(d3.axisBottom(xScale).ticks(8))
|
||||
.call(g => g.select('.domain').attr('stroke', '#ddd'))
|
||||
.call(g => g.selectAll('.tick line').remove())
|
||||
.call(g => g.selectAll('.tick text')
|
||||
.attr('font-size', '0.65rem')
|
||||
.attr('fill', '#888'))
|
||||
|
||||
// Create histogram bins
|
||||
const binGenerator = d3.bin()
|
||||
.domain(xScale.domain())
|
||||
.thresholds(30)
|
||||
|
||||
const bins = binGenerator(d.draws)
|
||||
const maxCount = d3.max(bins, b => b.length)
|
||||
|
||||
// Area generator
|
||||
const area = d3.area()
|
||||
.x(b => xScale((b.x0 + b.x1) / 2))
|
||||
.y0(yPos)
|
||||
.y1(yPos)
|
||||
.curve(d3.curveBasis)
|
||||
// X-axis label
|
||||
svg.append('text')
|
||||
.attr('x', margin.left + innerWidth / 2)
|
||||
.attr('y', height - 12)
|
||||
.attr('text-anchor', 'middle')
|
||||
.attr('font-size', '0.7rem')
|
||||
.attr('fill', '#888')
|
||||
.text('Arrests per 1,000 students')
|
||||
|
||||
// Scale heights
|
||||
const heightScale = d3.scaleLinear()
|
||||
.domain([0, maxCount])
|
||||
.range([rowHeight * 0.1, rowHeight * 0.9])
|
||||
|
||||
// Draw the ridgeline
|
||||
svg.append('path')
|
||||
.datum(bins)
|
||||
.attr('fill', colors.navy)
|
||||
.attr('opacity', 0.6)
|
||||
.attr('stroke', 'white')
|
||||
.attr('stroke-width', 0.5)
|
||||
.attr('d', d => area.bisector ?
|
||||
d3.area()
|
||||
.x(b => xScale((b.x0 + b.x1) / 2))
|
||||
.y0(yPos)
|
||||
.y1(b => yPos + heightScale(b.length))
|
||||
.curve(d3.curveBasis)(d)
|
||||
: null
|
||||
)
|
||||
|
||||
// Fix: manual polygon construction for each bin
|
||||
const points = []
|
||||
bins.forEach(b => {
|
||||
if (b.length > 0) {
|
||||
const xMid = (b.x0 + b.x1) / 2
|
||||
points.push([xScale(xMid), yPos])
|
||||
points.push([xScale(xMid), yPos + heightScale(b.length)])
|
||||
}
|
||||
})
|
||||
|
||||
// Add top edge
|
||||
for (let i = bins.length - 1; i >= 0; i--) {
|
||||
const b = bins[i]
|
||||
if (b.length > 0) {
|
||||
const xMid = (b.x0 + b.x1) / 2
|
||||
points.push([xScale(xMid), yPos])
|
||||
}
|
||||
}
|
||||
|
||||
if (points.length > 0) {
|
||||
svg.append('path')
|
||||
.attr('d', `M${points.map(p => p.join(',')).join('L')}Z`)
|
||||
.attr('fill', colors.navy)
|
||||
.attr('opacity', 0.55)
|
||||
.attr('stroke', 'white')
|
||||
.attr('stroke-width', 0.5)
|
||||
}
|
||||
|
||||
// Observed rate marker (diamond shape)
|
||||
if (d.observedRate > 0 && d.observedRate <= xMax) {
|
||||
const obsY = yPos + rowHeight / 2
|
||||
|
||||
// Diamond polygon
|
||||
svg.append('path')
|
||||
.attr('d', `M${xScale(d.observedRate)},${obsY - 6} L${xScale(d.observedRate) + 5},${obsY} L${xScale(d.observedRate)},${obsY + 6} L${xScale(d.observedRate) - 5},${obsY} Z`)
|
||||
.attr('fill', colors.danger)
|
||||
|
||||
// Label for observed rate
|
||||
svg.append('text')
|
||||
.attr('x', xScale(d.observedRate))
|
||||
.attr('y', obsY - 10)
|
||||
.attr('text-anchor', 'middle')
|
||||
.attr('font-size', '0.65rem')
|
||||
.attr('fill', colors.danger)
|
||||
.attr('font-weight', 600)
|
||||
.text(d.observedRate.toFixed(1))
|
||||
}
|
||||
})
|
||||
|
||||
// Y-axis labels
|
||||
// Y-axis
|
||||
svg.append('g')
|
||||
.attr('transform', `translate(${margin.left},0)`)
|
||||
.call(d3.axisLeft(yScale).tickSize(0))
|
||||
.call(g => g.select('.domain').remove())
|
||||
.call(g => g.selectAll('.tick text')
|
||||
.attr('font-size', '0.75rem')
|
||||
.attr('fill', colors.ink)
|
||||
.attr('font-weight', 500))
|
||||
|
||||
.attr('font-size', '0.65rem')
|
||||
.attr('fill', '#444'))
|
||||
|
||||
// Colors matching Civilytics palette
|
||||
const fillColor = '#000a9b' // navy
|
||||
const obsColor = '#c92d0e' // danger red
|
||||
|
||||
// Draw each ridge (one per race × sex group)
|
||||
plotData.forEach(d => {
|
||||
const yTop = yScale(d.label) || 0
|
||||
const bandwidth = yScale.bandwidth() || 40
|
||||
const baselineY = yTop + bandwidth * 0.15 // bottom of the ridge area
|
||||
const peakHeight = bandwidth * 0.8 // available height for density
|
||||
|
||||
// Create histogram bins from draws, then smooth into a ridge shape
|
||||
const nBins = 40
|
||||
const binWidth = maxRate / nBins
|
||||
const bins = new Array(nBins).fill(0)
|
||||
|
||||
d.draws.forEach(val => {
|
||||
if (val <= maxRate) {
|
||||
const idx = Math.min(Math.floor(val / binWidth), nBins - 1)
|
||||
if (idx >= 0) bins[idx]++
|
||||
}
|
||||
})
|
||||
|
||||
// Find peak count for scaling
|
||||
const maxCount = d3.max(bins) || 1
|
||||
const heightScale = peakHeight * 0.85 / maxCount
|
||||
|
||||
// Build the top edge points of the ridge (smoothed with curveBasis)
|
||||
const topPoints = bins.map((count, i) => {
|
||||
if (count === 0) return null
|
||||
const xMid = margin.left + ((i + 0.5) * binWidth / maxRate) * innerWidth
|
||||
// Ridge grows upward from baseline: higher count → taller ridge (lower y value)
|
||||
const yVal = baselineY - Math.min(count * heightScale, peakHeight)
|
||||
return [xMid, yVal]
|
||||
}).filter(Boolean)
|
||||
|
||||
if (topPoints.length < 2) return
|
||||
|
||||
// Smooth the top edge using basis interpolation
|
||||
const lineGen = d3.line()
|
||||
.curve(d3.curveBasis)
|
||||
.x(d => d[0])
|
||||
.y(d => d[1])
|
||||
|
||||
const smoothedTop = lineGen(topPoints) || ''
|
||||
|
||||
if (!smoothedTop) return
|
||||
|
||||
// Build the complete area path: bottom edge + smoothed top + close
|
||||
// The fill goes from baselineY down to peakHeight (upward in SVG coords, so y decreases)
|
||||
const firstX = margin.left
|
||||
const lastX = margin.left + innerWidth
|
||||
|
||||
let ridgePath = `M${firstX},${baselineY}` // start at bottom-left of ridge
|
||||
ridgePath += `L${topPoints[0][0]},${baselineY}` // line to first data x (along baseline)
|
||||
ridgePath += smoothedTop.substring(1) // append the smoothed top curve (skip 'M')
|
||||
ridgePath += `L${lastX},${baselineY}Z` // close back along bottom
|
||||
|
||||
svg.append('path')
|
||||
.attr('d', ridgePath)
|
||||
.attr('fill', fillColor)
|
||||
.attr('opacity', 0.55)
|
||||
.attr('stroke', 'white')
|
||||
.attr('stroke-width', 0.5)
|
||||
|
||||
// Observed rate marker — a diamond sitting on top of each ridge
|
||||
if (d.observedRate > 0 && d.observedRate <= maxRate) {
|
||||
const obsX = xScale(d.observedRate)
|
||||
const obsY = baselineY - peakHeight * 0.2 // position above the density
|
||||
|
||||
// Diamond shape (like ggplot2's shape=18)
|
||||
svg.append('polygon')
|
||||
.attr('points', [
|
||||
`${obsX},${obsY - 4}`,
|
||||
`${obsX + 4},${obsY}`,
|
||||
`${obsX},${obsY + 4}`,
|
||||
`${obsX - 4},${obsY}`
|
||||
].join(' '))
|
||||
.attr('fill', obsColor)
|
||||
|
||||
// Label the observed rate value below the diamond
|
||||
svg.append('text')
|
||||
.attr('x', obsX)
|
||||
.attr('y', yTop + bandwidth * 0.65)
|
||||
.attr('text-anchor', 'middle')
|
||||
.attr('font-size', '0.6rem')
|
||||
.attr('fill', '#888')
|
||||
.text(d.observedRate.toFixed(1))
|
||||
}
|
||||
})
|
||||
|
||||
// Legend
|
||||
const legend = svg.append('g')
|
||||
.attr('transform', `translate(${margin.left + innerWidth - 120}, ${height - 35})`)
|
||||
|
||||
// Diamond legend
|
||||
legend.append('path')
|
||||
.attr('d', 'M0,-4 L3,0 L0,4 L-3,0 Z')
|
||||
.attr('fill', colors.danger)
|
||||
legend.append('text')
|
||||
.attr('x', 8)
|
||||
.attr('y', 2)
|
||||
.attr('font-size', '0.65rem')
|
||||
.attr('fill', '#666')
|
||||
.text('Observed rate')
|
||||
|
||||
// Ridge legend
|
||||
legend.append('rect')
|
||||
.attr('x', 70)
|
||||
.attr('y', -8)
|
||||
const legendY = height - 25
|
||||
svg.append('rect')
|
||||
.attr('x', margin.left)
|
||||
.attr('y', legendY)
|
||||
.attr('width', 12)
|
||||
.attr('height', 8)
|
||||
.attr('fill', colors.navy)
|
||||
.attr('fill', fillColor)
|
||||
.attr('opacity', 0.55)
|
||||
legend.append('text')
|
||||
.attr('x', 86)
|
||||
.attr('y', -2)
|
||||
.attr('font-size', '0.65rem')
|
||||
.attr('fill', '#666')
|
||||
.text('Modeled (90%)')
|
||||
svg.append('text')
|
||||
.attr('x', margin.left + 16)
|
||||
.attr('y', legendY + 7)
|
||||
.attr('font-size', '0.62rem')
|
||||
.attr('fill', '#888')
|
||||
.text('Modeled posterior (90% interval)')
|
||||
|
||||
}, [quadData, selectedModel, rateByGroup])
|
||||
// Observed diamond legend
|
||||
svg.append('polygon')
|
||||
.attr('points', [
|
||||
`${margin.left + 140},${legendY}`,
|
||||
`${margin.left + 144},${legendY + 5}`,
|
||||
`${margin.left + 140},${legendY + 10}`,
|
||||
`${margin.left + 136},${legendY + 5}`
|
||||
].join(' '))
|
||||
.attr('fill', obsColor)
|
||||
svg.append('text')
|
||||
.attr('x', margin.left + 150)
|
||||
.attr('y', legendY + 7)
|
||||
.attr('font-size', '0.62rem')
|
||||
.attr('fill', '#888')
|
||||
.text('Observed rate (diamond)')
|
||||
|
||||
}, [quadData, selectedModel])
|
||||
|
||||
if (!quadData || Object.keys(quadData).length === 0) {
|
||||
return (
|
||||
@@ -288,8 +283,8 @@ export default function RateDensityRidgeline({ quadData, rateByGroup }) {
|
||||
</h3>
|
||||
|
||||
{/* Model selector */}
|
||||
<select
|
||||
value={selectedModel}
|
||||
<select
|
||||
value={selectedModel}
|
||||
onChange={(e) => setSelectedModel(e.target.value)}
|
||||
style={{ marginBottom: 'var(--space-2)', padding: '0.25rem 0.5rem', fontFamily: 'var(--font-sans)' }}
|
||||
>
|
||||
@@ -300,14 +295,13 @@ export default function RateDensityRidgeline({ quadData, rateByGroup }) {
|
||||
|
||||
{/* D3 Ridgeline */}
|
||||
<div style={{ width: '100%', overflowX: 'auto' }}>
|
||||
<svg ref={svgRef} style={{ width: '100%', minWidth: 400, height: 'auto' }} />
|
||||
<svg ref={svgRef} style={{ width: '100%', minWidth: 420, height: 'auto' }} />
|
||||
</div>
|
||||
|
||||
{/* Note */}
|
||||
<p style={{ fontSize: '0.72rem', color: 'var(--cv-ink-3)', marginTop: 'var(--space-1)' }}>
|
||||
<strong>Density ridges</strong> show the posterior distribution of predicted arrests per 1,000 students.
|
||||
<strong>Diamonds</strong> mark the observed rate for each group (the raw data point the model aims to improve).
|
||||
Using 90% intervals from the {MODEL_LABELS[selectedModel]} model.
|
||||
<p style={{ fontSize: '0.7rem', color: 'var(--cv-ink-3)', marginTop: 'var(--space-1)' }}>
|
||||
<strong>Density ridges</strong> show the posterior distribution of predicted arrests per 1,000 students.
|
||||
<strong>Diamonds</strong> mark observed rates for each group (the raw data point being modeled).
|
||||
</p>
|
||||
</div>
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user