| """Paper diagram for KVPath-Delta's low-rank path recurrence.""" | |
| from pathlib import Path | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| from matplotlib.patches import Circle, FancyArrowPatch, FancyBboxPatch | |
| OUT = Path(__file__).resolve().parent | |
| INK = "#242424" | |
| BLUE = "#36647D" | |
| ORANGE = "#A66B3F" | |
| PURPLE = "#72558C" | |
| BLUE_FILL = "#EDF3F6" | |
| ORANGE_FILL = "#F8EEE4" | |
| plt.rcParams.update({ | |
| "font.family": "STIXGeneral", | |
| "mathtext.fontset": "stix", | |
| "font.size": 11, | |
| "pdf.fonttype": 42, | |
| "svg.fonttype": "none", | |
| }) | |
| def text(ax, x, y, value, size=11, color=INK, **kwargs): | |
| ax.text(x, y, value, ha="center", va="center", fontsize=size, | |
| color=color, **kwargs) | |
| def arrow(ax, start, end, color=INK, lw=.9, mutation_scale=8): | |
| ax.add_patch(FancyArrowPatch( | |
| start, end, arrowstyle="-|>", mutation_scale=mutation_scale, | |
| linewidth=lw, color=color, | |
| )) | |
| fig, ax = plt.subplots(figsize=(8.4, 5.05)) | |
| ax.set(xlim=(.45, 11.35), ylim=(.08, 7.05)) | |
| ax.axis("off") | |
| state_x, state_w = 1.2, 5.2 | |
| delta_x, delta_w = 7.18, 1.40 | |
| input_x = 10.55 | |
| row_h = .78 | |
| centers = [1.35, 2.85, 4.35, 5.85] | |
| plus_x = state_x + state_w / 2 | |
| text(ax, plus_x, 6.72, "Cumulative full K/V", 12, BLUE) | |
| text(ax, delta_x + delta_w / 2, 6.72, "Low-rank K/V delta", 12, ORANGE) | |
| text(ax, delta_x + delta_w / 2, 6.45, | |
| r'$d_\delta\ll d_{KV}$', 9.5, ORANGE) | |
| text(ax, input_x, 6.72, "Layer input", 12) | |
| text(ax, .82, 6.72, "Layer", 10) | |
| for i, y in enumerate(centers): | |
| # There is one full K/V state, with no independent layer-specific branch. | |
| ax.add_patch(FancyBboxPatch( | |
| (state_x, y - row_h / 2), state_w, row_h, | |
| boxstyle="round,pad=0.015,rounding_size=.10", | |
| facecolor=BLUE_FILL, edgecolor=BLUE, linewidth=1.0, | |
| )) | |
| state_label = (r'$(k^\alpha,v^\alpha)$' if i == 0 | |
| else rf'$(k_{i},v_{i})$') | |
| text(ax, state_x + state_w / 2, y, state_label, 14, BLUE) | |
| text(ax, .82, y, rf'${i}$', 12) | |
| text(ax, input_x, y, rf'$h_{i}$', 15) | |
| if i == 0: | |
| # The anchor directly initializes a full-width K/V state. | |
| arrow(ax, (input_x - .38, y), (state_x + state_w + .08, y), BLUE) | |
| text(ax, 8.70, y + .20, r'$W^\alpha$ (anchor)', 11, BLUE) | |
| else: | |
| # Followers project only a compact delta. Its expansion enters an | |
| # explicit add node on the vertical state path, making the recurrence | |
| # visually read as "previous K/V + delta -> next K/V". | |
| ax.add_patch(FancyBboxPatch( | |
| (delta_x, y - .22), delta_w, .44, | |
| boxstyle="round,pad=0.012,rounding_size=.07", | |
| facecolor=ORANGE_FILL, edgecolor=ORANGE, linewidth=1.0, | |
| )) | |
| text(ax, delta_x + delta_w / 2, y, | |
| rf'$(k^\delta_{i},v^\delta_{i})$', 11.5, ORANGE) | |
| arrow(ax, (input_x - .38, y), (delta_x + delta_w + .06, y), ORANGE) | |
| text(ax, 9.20, y + .20, rf'$W^\delta_{i}$', 11, ORANGE) | |
| merge_y = (centers[i - 1] + y) / 2 | |
| elbow_x = state_x + state_w + .36 | |
| ax.plot([delta_x, elbow_x, elbow_x], [y, y, merge_y], | |
| color=ORANGE, linewidth=1.0) | |
| arrow(ax, (elbow_x, merge_y), (plus_x + .13, merge_y), | |
| ORANGE, lw=1.0, mutation_scale=9) | |
| text(ax, elbow_x + .10, merge_y + .20, | |
| rf'$U^K_{i},\ U^V_{i}$', 10.5, ORANGE) | |
| previous_top = centers[i - 1] + row_h / 2 + .03 | |
| current_bottom = y - row_h / 2 - .03 | |
| arrow(ax, (plus_x, previous_top), (plus_x, merge_y - .13), | |
| BLUE, lw=1.05, mutation_scale=9) | |
| ax.add_patch(Circle((plus_x, merge_y), .13, facecolor="white", | |
| edgecolor=ORANGE, linewidth=1.0, zorder=4)) | |
| text(ax, plus_x, merge_y - .005, r'$+$', 12, ORANGE, zorder=5) | |
| arrow(ax, (plus_x, merge_y + .13), (plus_x, current_bottom), | |
| BLUE, lw=1.05, mutation_scale=9) | |
| # The anchor is the alpha base. Followers produce separate compact K and V | |
| # deltas and add their independently expanded forms to the cumulative state. | |
| text(ax, 3.15, .30, | |
| r'$(k^\alpha,v^\alpha)=W^\alpha h_0$', 10.5, BLUE) | |
| text(ax, 7.45, .30, | |
| r'$k_i=k_{i-1}+U_i^Kk_i^\delta,\quad ' | |
| r'v_i=v_{i-1}+U_i^Vv_i^\delta$', | |
| 10.5, ORANGE) | |
| fig.subplots_adjust(left=.03, right=.985, top=.985, bottom=.025) | |
| for extension in ("png", "pdf", "svg"): | |
| target = OUT / f"kvpath_delta_analogy.{extension}" | |
| fig.savefig(target, dpi=240, bbox_inches="tight", pad_inches=.035, | |
| facecolor="white") | |
| print(target) | |
| plt.close(fig) | |
Xet Storage Details
- Size:
- 4.54 kB
- Xet hash:
- d19338501c96c69150445cd2721f982aae020951f1361f579c19a31d6554e922
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.