xuan-luo/temp / docs /figures /plot_kvpath_variants.py
xuan-luo's picture
download
raw
5.51 kB
"""Three paper figure alternatives; all show a single four-layer path."""
from pathlib import Path
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
from matplotlib.patches import Rectangle, FancyArrowPatch
from matplotlib.path import Path as MplPath
OUT = Path(__file__).resolve().parent
plt.rcParams.update({'font.family': 'STIXGeneral', 'mathtext.fontset': 'stix',
'font.size': 11, 'pdf.fonttype': 42, 'svg.fonttype': 'none'})
INK, BLUE, ORANGE = '#242424', '#36647D', '#A66B3F'
def text(ax, x, y, s, size=11, color=INK, **kw):
ax.text(x, y, s, ha='center', va='center', fontsize=size, color=color, **kw)
def arrow(ax, a, b, color=INK, dashed=False):
ax.add_patch(FancyArrowPatch(a, b, arrowstyle='-|>', mutation_scale=8,
lw=.85, color=color, linestyle='--' if dashed else '-'))
def line(ax, xs, ys, color=INK, **kw):
ax.plot(xs, ys, color=color, lw=.85, **kw)
def structure(ax):
# The large alpha cell spans four contiguous delta cells.
ax.add_patch(Rectangle((1.2, 1), 2.4, 4.8,
facecolor='#EDF3F6', edgecolor=BLUE, lw=.9))
text(ax, 2.4, 3.65, r'$(k^\alpha,v^\alpha)$', 17, BLUE)
text(ax, 2.4, 3.13, 'shared across layers', 10, BLUE)
text(ax, 2.4, 6.26, 'Alpha component', 12, BLUE)
text(ax, 4.6, 6.26, 'Delta component', 12, ORANGE)
text(ax, .82, 6.26, 'Layer', 10)
for i in range(4):
y = 1.6+1.2*i
ax.add_patch(Rectangle((3.6, y-.6), 2, 1.2,
facecolor='#F8EEE4', edgecolor=ORANGE, lw=.9))
text(ax, 4.6, y, rf'$(k^\delta_{i},v^\delta_{i})$', 14, ORANGE)
text(ax, .82, y, rf'${i}$', 12)
line(ax, [.4, .25, .25, .4], [1, 1, 5.8, 5.8])
text(ax, -.05, 3.4, r'One path ($m=4$)', 11, rotation=90)
def inputs(ax, analogy=False):
text(ax, 7.5, 6.26, 'Layer input', 12)
for i in range(4):
y = 1.6+1.2*i
text(ax, 7.5, y, rf'$h_{i}$', 15)
arrow(ax, (7.15, y), (5.65, y), ORANGE)
text(ax, 6.42, y+.23, rf'$W^\delta_{i}$', 12, ORANGE)
if analogy and i < 3:
arrow(ax, (7.5, y+.24), (7.5, y+.96), INK, dashed=True)
text(ax, 8.04, y+.60, rf'$\Delta h_{i+1}$', 12)
line(ax, [7.5, 7.5, 2.4], [1.35, .5, .5], BLUE)
arrow(ax, (2.4, .5), (2.4, .98), BLUE)
text(ax, 4.6, .73, r'$W^\alpha$', 12, BLUE)
def save(fig, name):
fig.subplots_adjust(left=.03, right=.98, top=.98, bottom=.025)
for ext in ('png', 'pdf', 'svg'):
path = OUT / f'kvpath_{name}.{ext}'
fig.savefig(path, dpi=240, bbox_inches='tight', pad_inches=.035,
facecolor='white')
print(path)
plt.close(fig)
# A: just the implemented construction; no cross-layer difference arrows.
fig, ax = plt.subplots(figsize=(6.5, 4.35))
ax.set(xlim=(-.4, 8.25), ylim=(.30, 6.65)); ax.axis('off')
structure(ax); inputs(ax)
save(fig, 'a_structure')
# B: the cache layout and one adjacent-row difference, annotated outside
# the cells so there is no implied recurrent computation or artificial gap.
fig, ax = plt.subplots(figsize=(6.5, 4.35))
ax.set(xlim=(-.4, 9.5), ylim=(.30, 6.65)); ax.axis('off')
structure(ax)
line(ax, [5.74, 5.99, 5.99, 5.74], [2.8, 2.8, 4, 4], ORANGE)
line(ax, [5.99, 6.36], [3.4, 3.4], ORANGE)
text(ax, 7.8, 4.05, 'Adjacent-layer difference', 11)
text(ax, 7.8, 3.48, r'$\Delta k_2=[\,0\,\Vert\,k^\delta_2-k^\delta_1\,]$', 13)
text(ax, 7.8, 2.89, r'$\Delta v_2=[\,0\,\Vert\,v^\delta_2-v^\delta_1\,]$', 13)
text(ax, 2.4, .63, r'$d_\alpha$ dimensions', 11, BLUE)
text(ax, 4.6, .63, r'$d_\delta$ dimensions', 11, ORANGE)
save(fig, 'b_difference')
# C: differences in h and KV are juxtaposed, never connected by a map.
# h is a normalized layer input, so Delta h denotes a difference, not a
# literal Transformer residual branch.
fig, ax = plt.subplots(figsize=(8.0, 4.5))
ax.set(xlim=(.45, 10.85), ylim=(.30, 6.55)); ax.axis('off')
def difference_arrow(x, outer_x, low, high, color):
"""Right-angle dashed path: right, up, then a left-pointing arrow."""
path = MplPath([(x, low), (outer_x, low), (outer_x, high), (x, high)],
[MplPath.MOVETO, MplPath.LINETO, MplPath.LINETO, MplPath.LINETO])
ax.add_patch(FancyArrowPatch(path=path, arrowstyle='-|>',
mutation_scale=8, linewidth=.85, linestyle=(0, (3.2, 2.2)),
color=color, joinstyle='miter', capstyle='butt'))
# Each layer stores one concatenated [alpha | delta] block. The four blocks
# are separated vertically so the upward inter-layer links remain visible.
text(ax, 3.2, 6.22, 'Alpha component', 12, BLUE)
text(ax, 6.2, 6.22, 'Delta component', 12, ORANGE)
text(ax, 10.15, 6.22, 'Layer input', 12)
text(ax, .82, 6.22, 'Layer', 10)
for i in range(4):
y = 1.45+1.3*i
ax.add_patch(Rectangle((1.2, y-.5), 4, 1,
facecolor='#EDF3F6', edgecolor=BLUE, lw=.9))
text(ax, 3.2, y, r'$(k^\alpha,v^\alpha)$', 14, BLUE)
ax.add_patch(Rectangle((5.2, y-.5), 2, 1,
facecolor='#F8EEE4', edgecolor=ORANGE, lw=.9))
text(ax, 6.2, y, rf'$(k^\delta_{i},v^\delta_{i})$', 14, ORANGE)
text(ax, .82, y, rf'${i}$', 12)
text(ax, 10.15, y, rf'$h_{i}$', 15)
arrow(ax, (9.78, y), (7.25, y), ORANGE)
text(ax, 8.75, y+.22, rf'$W^\delta_{i}$', 12, ORANGE)
if i < 3:
arrow(ax, (10.15, y+.25), (10.15, y+.95), INK)
arrow(ax, (5.2, y+.52), (5.2, y+.78), ORANGE)
line(ax, [10.15, 10.15, 3.2], [1.20, .5, .5], BLUE)
arrow(ax, (3.2, .5), (3.2, .93), BLUE)
text(ax, 6.2, .73, r'$W^\alpha$', 12, BLUE)
save(fig, 'c_analogy')

Xet Storage Details

Size:
5.51 kB
·
Xet hash:
ab0f9f8046c4621f6773aa4ea671228e7d451fe9d5af285ee57c44e8b646e409

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.