xuan-luo/temp / docs /figures /plot_kvpath_delta.py
xuan-luo's picture
download
raw
4.54 kB
"""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.