| """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.