# Setup, inlined rather than imported so this notebook runs on its own.
# Keep it folded; nothing below it depends on anything outside this file.
import matplotlib.pyplot as plt
# --- chart styling -------------------------------------------------------
# Categorical slots of a CVD-validated palette: blue, orange, aqua, purple.
COLOURS = ["#2a78d6", "#eb6834", "#1baf7a", "#8a63d2"]
MUTED, GRID, AXIS = "#5b6570", "#e6e6e3", "#d5d5d1"
def style_axes(ax, xlabel=None, ylabel=None, grid="y"):
"""Strip an axes back to the ink that carries information."""
if xlabel:
ax.set_xlabel(xlabel, color=MUTED, fontsize=9)
if ylabel:
ax.set_ylabel(ylabel, color=MUTED, fontsize=9)
if grid:
ax.grid(axis=grid, color=GRID, linewidth=0.8)
ax.set_axisbelow(True)
for side in ("top", "right"):
ax.spines[side].set_visible(False)
for side in ("left", "bottom"):
ax.spines[side].set_color(AXIS)
ax.tick_params(colors=MUTED, labelsize=9, length=0)
return ax
def figure(width=7.0, height=4.2, **kw):
fig, ax = plt.subplots(figsize=(width, height), **kw)
return fig, ax
# --- the content-scoring first layer, drawn to scale ----------------------
# Late fusion splits the layer in proportion to how many columns each group
# brings, so the title takes 233 of 256 units and the author's three scalars get
# 2. A grid of equal cells hides exactly that, which is why this is drawn.
FUSION_GROUPS = [("title", 300, 233), ("title shape", 11, 9), ("where it links", 9, 7),
("when posted", 6, 5), ("author history", 3, 2)]
_IN_X, _IN_W, _HID_X, _HID_W = 0.30, 0.055, 0.70, 0.055
_TOP, _BOTTOM, _GAP = 0.90, 0.12, 0.006
def _stack(ax, values, total, x, width, colours):
span = (_TOP - _BOTTOM) - _GAP * (len(values) - 1)
y, out = _BOTTOM, []
for value, colour in zip(values, colours):
height = span * value / total
ax.add_patch(plt.Rectangle((x, y), width, height, facecolor=colour,
edgecolor="none"))
out.append((y, height))
y += height + _GAP
return out
def _wedge(ax, a, b, alpha):
(y0, h0), (y1, h1) = a, b
ax.fill([_IN_X + _IN_W, _IN_X + _IN_W, _HID_X, _HID_X],
[y0, y0 + h0, y1 + h1, y1], color=COLOURS[2], alpha=alpha, lw=0, zorder=0)
def _bracket(ax, x, low, high, label, align):
ax.plot([x] * 2, [low, high], color="#c3ccd6", linewidth=1)
ax.text(x + (-0.007 if align == "right" else 0.007), (low + high) / 2, label,
va="center", ha=align, fontsize=8, color=MUTED, linespacing=1.4)
def fusion_layer(ax, kind, title):
"""One panel: the 329 inputs, and the 256 hidden units they may reach."""
ax.set_xlim(0, 1)
ax.set_ylim(0, 1)
ax.axis("off")
ax.set_title(title, fontsize=10, color=MUTED, pad=6)
inputs = _stack(ax, [g[1] for g in FUSION_GROUPS], 329, _IN_X, _IN_W,
[COLOURS[0]] + ["#9fb6cf"] * 4)
if kind == "early":
hidden = _stack(ax, [256], 256, _HID_X, _HID_W, [COLOURS[2]])
for group in inputs:
_wedge(ax, group, hidden[0], 0.16)
ax.text(_HID_X + _HID_W + 0.025, hidden[0][0] + hidden[0][1] / 2, "256 units",
va="center", fontsize=9, color=MUTED)
else:
hidden = _stack(ax, [g[2] for g in FUSION_GROUPS], 256, _HID_X, _HID_W,
[COLOURS[2]] + ["#9ed4bd"] * 4)
for group, block_ in zip(inputs, hidden):
_wedge(ax, group, block_, 0.30)
ax.text(_HID_X + _HID_W + 0.025, hidden[0][0] + hidden[0][1] / 2, "233 units",
va="center", fontsize=9, color=MUTED)
_bracket(ax, _HID_X + _HID_W + 0.018, hidden[1][0],
hidden[-1][0] + hidden[-1][1], "9, 7, 5, 2\nbetween them", "left")
ax.text(_IN_X - 0.025, inputs[0][0] + inputs[0][1] / 2, "title 300",
va="center", ha="right", fontsize=9, color=MUTED)
_bracket(ax, _IN_X - 0.018, inputs[1][0], inputs[-1][0] + inputs[-1][1],
"four metadata groups\n11, 9, 6, 3 columns", "right")
ax.text(_IN_X + _IN_W / 2, 0.045, "329 inputs", fontsize=8.5, color=MUTED,
ha="center")
ax.text(_HID_X + _HID_W / 2, 0.045, "256 hidden units", fontsize=8.5, color=MUTED,
ha="center")
import json
import numpy as np
from huggingface_hub import hf_hub_download
REVISION = "f7f04575591e5dfb4996534911c0c55e993f3899"
def results(name):
"""The raw per-seed measurements, pinned to a dataset commit."""
return json.load(open(hf_hub_download("roshbeed/ai-residency-blog-data",
f"fusion/{name}.json",
repo_type="dataset", revision=REVISION)))
synthetic = results("synthetic-alpha")
counts = results("pairings-counts")
logged = results("pairings-log")
ranked = results("pairings-rank")
replicate = results("pairings-replicate")
robust_128 = results("robustness-128")
robust_512 = results("robustness-512")
print(f"synthetic sweep: {len(synthetic)} interaction strengths x "
f"{len(next(iter(synthetic.values())))} model variants")
print(f"Hacker News: {len(counts)} pairings of input groups, 10 seeds each,")
print(f" measured on raw counts, on log1p(counts) and on rank")