tnh0527's picture
Publish ettin-150m-memory-reranker-ft-v1 (qualified ONNX export and complete attribution)
bf518ed verified
Raw History Blame Contribute Delete
21.5 kB
"""Dependency-free SVG charts for tracked evaluation figures.
The evaluation figures are generated from private per-row evidence into tracked SVG files.
Keeping the renderer inside the repository, with no plotting dependency, means a figure can be
regenerated by any contributor with the private inputs and compared byte for byte.
Everything is emitted as presentation attributes rather than CSS. Markdown hosts sanitize
embedded stylesheets and scripts out of SVG, so a figure that carries its styling in attributes
renders the same in the repository, on a model-card host, and in a local viewer.
Layout is measured rather than assumed: legend entries, panel titles, and reference-line labels
are placed from an estimated text width, so a longer label reflows instead of overlapping its
neighbour. ``text_width`` approximates a sans-serif advance table, which is enough to keep
elements apart but is not a substitute for a real font metric.
"""
from __future__ import annotations
import math
from dataclasses import dataclass, field
from xml.sax.saxutils import escape
# Okabe-Ito, chosen because it stays distinguishable under the common colour-vision
# deficiencies and prints legibly in greyscale.
PALETTE = ("#0072B2", "#D55E00", "#009E73", "#CC79A7", "#E69F00", "#56B4E9", "#000000")
FONT_STACK = "system-ui, Segoe UI, Roboto, Helvetica, Arial, sans-serif"
FONT = f'font-family="{FONT_STACK}"'
INK = "#1a1a1a"
MUTED = "#5c5c5c"
AXIS = "#8a8a8a"
GRID = "#e8e8e8"
GRID_STRONG = "#d0d0d0"
# Per-character advance as a fraction of font size, for a humanist sans at normal weight.
_NARROW = set("iljft.,;:|!()[]{}I '")
_WIDE = set("mwMW@%")
_DIGIT = set("0123456789")
def text_width(text: str, size: float, *, weight: str = "normal") -> float:
"""Estimate rendered width in user units."""
total = 0.0
for character in text:
if character in _NARROW:
total += 0.30
elif character in _WIDE:
total += 0.86
elif character in _DIGIT:
total += 0.56
elif character.isupper():
total += 0.66
else:
total += 0.52
if weight in {"600", "700", "bold"}:
total *= 1.05
return total * size
def _fit_lines(text: str, size: float, limit: float, *, weight: str = "normal") -> list[str]:
"""Wrap to at most two lines, breaking on whitespace."""
if text_width(text, size, weight=weight) <= limit:
return [text]
words = text.split(" ")
line: list[str] = []
for index, word in enumerate(words):
candidate = " ".join([*line, word])
if line and text_width(candidate, size, weight=weight) > limit:
return [" ".join(line), " ".join(words[index:])]
line.append(word)
return [" ".join(line)]
@dataclass
class Series:
label: str
x: list[float]
y: list[float]
color: str = PALETTE[0]
dash: str | None = None
marker: bool = True
width: float = 1.9
marker_size: float = 2.6
@dataclass
class Band:
"""A shaded interval drawn behind its series."""
x: list[float]
low: list[float]
high: list[float]
color: str = PALETTE[0]
opacity: float = 0.16
@dataclass
class Counts:
"""A support strip under the plot: how many rows sit behind each x position."""
x: list[float]
values: list[float]
color: str = MUTED
label: str = "rows per bin"
@dataclass
class Panel:
title: str
series: list[Series] = field(default_factory=list)
bands: list[Band] = field(default_factory=list)
xlabel: str = ""
ylabel: str = ""
xlim: tuple[float, float] | None = None
ylim: tuple[float, float] | None = None
xticks: list[tuple[float, str]] | None = None
yticks: list[tuple[float, str]] | None = None
xscale: str = "linear"
hlines: list[tuple[float, str, str]] = field(default_factory=list)
diagonal: bool = False
notes: list[str] = field(default_factory=list)
counts: Counts | None = None
def _fmt(value: float) -> str:
if abs(value) >= 1e6:
return f"{value:.3g}"
text = f"{value:.3f}".rstrip("0").rstrip(".")
return text or "0"
def _auto_ticks(low: float, high: float, count: int = 5) -> list[tuple[float, str]]:
if high <= low:
high = low + 1.0
step = (high - low) / count
magnitude = 10 ** math.floor(math.log10(step)) if step > 0 else 1.0
for factor in (1, 2, 2.5, 5, 10):
if step <= factor * magnitude:
step = factor * magnitude
break
start = math.ceil(low / step) * step
ticks = []
value = start
while value <= high + 1e-9:
ticks.append((value, _fmt(0.0 if abs(value) < step * 1e-6 else value)))
value += step
return ticks
def _text(
x: float,
y: float,
body: str,
*,
size: float,
fill: str = INK,
anchor: str = "start",
weight: str | None = None,
) -> str:
weight_attr = f' font-weight="{weight}"' if weight else ""
return (
f'<text x="{x:.1f}" y="{y:.1f}" text-anchor="{anchor}" font-size="{size}" '
f'fill="{fill}"{weight_attr} {FONT}>{escape(body)}</text>'
)
_TITLE_SIZE = 12.5
# Everything below the plot box is stacked in fixed bands rather than placed at absolute
# offsets, so an x-axis label, a support strip, and a note can coexist without overlapping.
_TICK_BAND = 18.0
_XLABEL_BAND = 16.0
_COUNTS_BAND = 30.0
_NOTE_BAND = 12.0
_NOTE_LEAD = 6.0
_FLOOR_SLACK = 8.0
def _panel_bottom(panel: Panel) -> float:
bottom = _TICK_BAND + _FLOOR_SLACK
if panel.xlabel:
bottom += _XLABEL_BAND
if panel.counts:
bottom += _COUNTS_BAND
if panel.notes:
bottom += _NOTE_LEAD + _NOTE_BAND * len(panel.notes)
return bottom
def _panel_svg(panel: Panel, width: float, height: float, *, title_rows: int | None = None) -> str:
left, right = 54.0, 14.0
plot_w = width - left - right
title_lines = _fit_lines(panel.title, _TITLE_SIZE, plot_w, weight="600")
top = 18.0 + 14.0 * (title_rows or len(title_lines))
plot_h = height - top - _panel_bottom(panel)
xs = [x for s in panel.series for x in s.x] or [0.0, 1.0]
ys = [y for s in panel.series for y in s.y] or [0.0, 1.0]
ys += [value for band in panel.bands for value in (*band.low, *band.high)]
ys += [level for level, _, _ in panel.hlines]
xlim = panel.xlim or (min(xs), max(xs))
ylim = panel.ylim or (min(ys), max(ys))
if ylim[0] == ylim[1]:
ylim = (ylim[0] - 0.5, ylim[1] + 0.5)
def tx(value: float) -> float:
if panel.xscale == "log":
lo, hi = math.log(xlim[0]), math.log(xlim[1])
return left + (math.log(value) - lo) / (hi - lo) * plot_w
return left + (value - xlim[0]) / (xlim[1] - xlim[0]) * plot_w
def ty(value: float) -> float:
return top + plot_h - (value - ylim[0]) / (ylim[1] - ylim[0]) * plot_h
parts: list[str] = []
for index, line in enumerate(title_lines):
parts.append(
_text(
left + plot_w / 2,
16.0 + 14.0 * index,
line,
size=_TITLE_SIZE,
anchor="middle",
weight="600",
)
)
xticks = panel.xticks or _auto_ticks(*xlim)
yticks = panel.yticks or _auto_ticks(*ylim)
for value, label in yticks:
if not ylim[0] - 1e-9 <= value <= ylim[1] + 1e-9:
continue
y = ty(value)
parts.append(
f'<line x1="{left}" y1="{y:.1f}" x2="{left + plot_w:.1f}" y2="{y:.1f}" '
f'stroke="{GRID}" stroke-width="1"/>'
)
parts.append(_text(left - 6, y + 3.4, label, size=10, fill=MUTED, anchor="end"))
for value, label in xticks:
if not xlim[0] - 1e-9 <= value <= xlim[1] + 1e-9:
continue
x = tx(value)
parts.append(
f'<line x1="{x:.1f}" y1="{top}" x2="{x:.1f}" y2="{top + plot_h:.1f}" '
f'stroke="{GRID}" stroke-width="1"/>'
)
parts.append(_text(x, top + plot_h + 13, label, size=10, fill=MUTED, anchor="middle"))
if panel.diagonal:
parts.append(
f'<line x1="{tx(xlim[0]):.1f}" y1="{ty(ylim[0]):.1f}" '
f'x2="{tx(xlim[1]):.1f}" y2="{ty(ylim[1]):.1f}" stroke="{AXIS}" '
f'stroke-width="1.1" stroke-dasharray="4 3"/>'
)
for band in panel.bands:
forward = " ".join(
f"{tx(x):.1f},{ty(y):.1f}" for x, y in zip(band.x, band.high, strict=True)
)
backward = " ".join(
f"{tx(x):.1f},{ty(y):.1f}"
for x, y in zip(reversed(band.x), reversed(band.low), strict=True)
)
parts.append(
f'<polygon points="{forward} {backward}" fill="{band.color}" '
f'fill-opacity="{band.opacity}" stroke="none"/>'
)
# A reference line carries its label at the left margin over a solid backing box, so the
# label never lands on the data it is a reference for.
for value, label, color in panel.hlines:
y = ty(value)
parts.append(
f'<line x1="{left}" y1="{y:.1f}" x2="{left + plot_w:.1f}" y2="{y:.1f}" '
f'stroke="{color}" stroke-width="1.2" stroke-dasharray="5 3"/>'
)
if not label:
continue
label_w = text_width(label, 9.5) + 8.0
parts.append(
f'<rect x="{left + 3:.1f}" y="{y - 12:.1f}" width="{label_w:.1f}" height="11.5" '
f'fill="white" fill-opacity="0.9" stroke="none"/>'
)
parts.append(_text(left + 7, y - 3.5, label, size=9.5, fill=color))
for series in panel.series:
points = " ".join(
f"{tx(x):.1f},{ty(y):.1f}" for x, y in zip(series.x, series.y, strict=True)
)
dash = f' stroke-dasharray="{series.dash}"' if series.dash else ""
parts.append(
f'<polyline points="{points}" fill="none" stroke="{series.color}" '
f'stroke-width="{series.width}" stroke-linejoin="round" '
f'stroke-linecap="round"{dash}/>'
)
if series.marker:
for x, y in zip(series.x, series.y, strict=True):
parts.append(
f'<circle cx="{tx(x):.1f}" cy="{ty(y):.1f}" r="{series.marker_size}" '
f'fill="{series.color}"/>'
)
parts.append(
f'<rect x="{left}" y="{top}" width="{plot_w:.1f}" height="{plot_h:.1f}" fill="none" '
f'stroke="{AXIS}" stroke-width="1"/>'
)
cursor = top + plot_h + _TICK_BAND
if panel.xlabel:
parts.append(
_text(
left + plot_w / 2,
cursor + 11,
panel.xlabel,
size=10.5,
fill=MUTED,
anchor="middle",
)
)
cursor += _XLABEL_BAND
if panel.counts:
counts = panel.counts
if not counts.values:
raise ValueError("a support strip needs at least one count")
strip_top = cursor + 2.0
strip_h = 15.0
peak = max(counts.values) or 1.0
slot = plot_w / max(len(counts.x), 1) * 0.7
for x, value in zip(counts.x, counts.values, strict=True):
bar_h = (value / peak) * strip_h
parts.append(
f'<rect x="{tx(x) - slot / 2:.1f}" y="{strip_top + strip_h - bar_h:.1f}" '
f'width="{slot:.1f}" height="{bar_h:.1f}" fill="{counts.color}" '
f'fill-opacity="0.5"/>'
)
parts.append(
f'<line x1="{left}" y1="{strip_top + strip_h:.1f}" x2="{left + plot_w:.1f}" '
f'y2="{strip_top + strip_h:.1f}" stroke="{GRID_STRONG}" stroke-width="1"/>'
)
parts.append(_text(left, strip_top + strip_h + 9, counts.label, size=9, fill=MUTED))
parts.append(
_text(
left + plot_w,
strip_top + strip_h + 9,
f"tallest {int(peak):,}",
size=9,
fill=MUTED,
anchor="end",
)
)
cursor += _COUNTS_BAND
for index, note in enumerate(panel.notes):
parts.append(
_text(left, cursor + _NOTE_LEAD + 9 + _NOTE_BAND * index, note, size=9.5, fill=MUTED)
)
if panel.ylabel:
parts.append(
f'<text transform="translate(13,{top + plot_h / 2:.1f}) rotate(-90)" '
f'text-anchor="middle" font-size="10.5" fill="{MUTED}" {FONT}>'
f"{escape(panel.ylabel)}</text>"
)
return "\n".join(parts)
def _legend_svg(
entries: list[tuple[str, str, str | None]],
*,
x: float,
y: float,
max_width: float,
) -> tuple[str, float]:
"""Flow legend entries across as many rows as their measured widths need."""
swatch, gap, pad = 22.0, 7.0, 22.0
parts: list[str] = []
cursor_x, cursor_y, rows = x, y, 1
for label, color, dash in entries:
entry_w = swatch + gap + text_width(label, 11) + pad
if cursor_x > x and cursor_x + entry_w - pad > x + max_width:
cursor_x, cursor_y, rows = x, cursor_y + 16.0, rows + 1
dash_attr = f' stroke-dasharray="{dash}"' if dash else ""
parts.append(
f'<line x1="{cursor_x:.1f}" y1="{cursor_y:.1f}" x2="{cursor_x + swatch:.1f}" '
f'y2="{cursor_y:.1f}" stroke="{color}" stroke-width="2.4" '
f'stroke-linecap="round"{dash_attr}/>'
)
parts.append(_text(cursor_x + swatch + gap, cursor_y + 3.8, label, size=11))
cursor_x += entry_w
return "\n".join(parts), 16.0 * rows
def _dedupe(entries: list[tuple[str, str, str | None]]) -> list[tuple[str, str, str | None]]:
seen: set[tuple[str, str, str | None]] = set()
unique = []
for entry in entries:
if entry not in seen:
seen.add(entry)
unique.append(entry)
return unique
def _open_svg(width: float, height: float, title: str, description: str) -> list[str]:
return [
f'<svg xmlns="http://www.w3.org/2000/svg" width="{width:.0f}" height="{height:.0f}" '
f'viewBox="0 0 {width:.0f} {height:.0f}" role="img" aria-label="{escape(title)}">',
f"<title>{escape(title)}</title>",
f"<desc>{escape(description)}</desc>",
'<rect width="100%" height="100%" fill="white"/>',
]
def render_grid(
panels: list[Panel],
*,
columns: int,
title: str,
subtitle: str = "",
panel_width: float = 352.0,
panel_height: float = 256.0,
legend: list[tuple[str, str, str | None]] | None = None,
description: str = "",
) -> str:
"""Render panels on a grid with one shared title and a flowed legend.
A final short row is centred, so a five-panel figure on three columns has no empty cell.
"""
if not panels:
raise ValueError("render_grid needs at least one panel")
if columns < 1:
raise ValueError("render_grid needs at least one column")
legend = _dedupe(legend or [])
rows = math.ceil(len(panels) / columns)
width = columns * panel_width
header = 26.0 + (16.0 if subtitle else 0.0)
legend_svg, legend_h = "", 0.0
if legend:
legend_svg, legend_h = _legend_svg(legend, x=18.0, y=header + 10.0, max_width=width - 36.0)
legend_h += 8.0
height = header + legend_h + rows * panel_height
parts = _open_svg(width, height, title, description or title)
parts.append(_text(width / 2, 19, title, size=14.5, anchor="middle", weight="700"))
if subtitle:
parts.append(_text(width / 2, 34, subtitle, size=11, fill=MUTED, anchor="middle"))
if legend_svg:
parts.append(legend_svg)
# One title row count for the whole grid, so a panel whose title wraps does not push its
# plot box below its neighbours'.
title_rows = max(
len(_fit_lines(panel.title, _TITLE_SIZE, panel_width - 68.0, weight="600"))
for panel in panels
)
for index, panel in enumerate(panels):
row, column = divmod(index, columns)
in_row = min(columns, len(panels) - row * columns)
offset = (columns - in_row) * panel_width / 2.0
px = offset + column * panel_width
py = header + legend_h + row * panel_height
parts.append(f'<g transform="translate({px:.1f},{py:.1f})">')
parts.append(_panel_svg(panel, panel_width, panel_height, title_rows=title_rows))
parts.append("</g>")
parts.append("</svg>")
return "\n".join(parts) + "\n"
def render_bars(
groups: list[str],
series: list[tuple[str, list[float], str]],
*,
title: str,
ylabel: str,
subtitle: str = "",
ylim: tuple[float, float] = (0.0, 1.0),
reference: list[tuple[str, list[float], str]] | None = None,
notes: list[str] | None = None,
separator_before: int | None = None,
width: float = 780.0,
height: float = 350.0,
description: str = "",
) -> str:
"""Render grouped bars with per-group dashed reference levels.
Bars keep a zero baseline. Value labels are drawn only where a bar is wide enough to hold
one, because a crowded label is worse than none.
"""
if not groups or not series:
raise ValueError("render_bars needs at least one group and one series")
if any(len(values) != len(groups) for _, values, _ in series):
raise ValueError("every bar series must carry one value per group")
notes = notes or []
left, right, top = 54.0, 16.0, 26.0 + (15.0 if subtitle else 0.0)
legend_entries = [(label, color, None) for label, _, color in series]
legend_entries += [(label, color, "5 3") for label, _, color in reference or []]
legend_svg, legend_h = _legend_svg(
_dedupe(legend_entries), x=18.0, y=top + 12.0, max_width=width - 36.0
)
top += legend_h + 10.0
bottom = 46.0 + 12.0 * len(notes)
plot_w, plot_h = width - left - right, height - top - bottom
group_w = plot_w / len(groups)
bar_w = group_w * 0.74 / len(series)
def ty(value: float) -> float:
return top + plot_h - (value - ylim[0]) / (ylim[1] - ylim[0]) * plot_h
parts = _open_svg(width, height, title, description or title)
parts.append(_text(width / 2, 19, title, size=14.5, anchor="middle", weight="700"))
if subtitle:
parts.append(_text(width / 2, 34, subtitle, size=11, fill=MUTED, anchor="middle"))
parts.append(legend_svg)
for value, label in _auto_ticks(*ylim):
y = ty(value)
parts.append(
f'<line x1="{left}" y1="{y:.1f}" x2="{left + plot_w:.1f}" y2="{y:.1f}" '
f'stroke="{GRID}" stroke-width="1"/>'
)
parts.append(_text(left - 6, y + 3.4, label, size=10, fill=MUTED, anchor="end"))
label_fits = bar_w - 2 >= text_width("0.000", 8.5) + 2
for g_index, group in enumerate(groups):
gx = left + g_index * group_w + group_w * 0.13
for s_index, (_, values, color) in enumerate(series):
value = values[g_index]
x = gx + s_index * bar_w
parts.append(
f'<rect x="{x:.1f}" y="{ty(value):.1f}" width="{bar_w - 2:.1f}" '
f'height="{ty(ylim[0]) - ty(value):.1f}" fill="{color}"/>'
)
if label_fits:
parts.append(
_text(
x + (bar_w - 2) / 2,
ty(value) - 4,
f"{value:.3f}",
size=8.5,
anchor="middle",
)
)
for _, values, color in reference or []:
y = ty(values[g_index])
parts.append(
f'<line x1="{gx - 3:.1f}" y1="{y:.1f}" '
f'x2="{gx + group_w * 0.74 + 3:.1f}" y2="{y:.1f}" stroke="{color}" '
f'stroke-width="1.6" stroke-dasharray="5 3"/>'
)
parts.append(
_text(
left + g_index * group_w + group_w / 2,
top + plot_h + 16,
group,
size=11,
anchor="middle",
)
)
if separator_before is not None and 0 < separator_before < len(groups):
x = left + separator_before * group_w
parts.append(
f'<line x1="{x:.1f}" y1="{top}" x2="{x:.1f}" y2="{top + plot_h + 6:.1f}" '
f'stroke="{GRID_STRONG}" stroke-width="1.4"/>'
)
parts.append(
f'<rect x="{left}" y="{top}" width="{plot_w:.1f}" height="{plot_h:.1f}" fill="none" '
f'stroke="{AXIS}" stroke-width="1"/>'
)
parts.append(
f'<text transform="translate(13,{top + plot_h / 2:.1f}) rotate(-90)" '
f'text-anchor="middle" font-size="10.5" fill="{MUTED}" {FONT}>'
f"{escape(ylabel)}</text>"
)
for index, note in enumerate(notes):
parts.append(_text(left, top + plot_h + 34 + 12 * index, note, size=9.5, fill=MUTED))
parts.append("</svg>")
return "\n".join(parts) + "\n"