GroundFlow / visualize /plot_attention_masks.py
TerryPei's picture
sync visualize from /opt/tiger/thothvl_pretrain (HF tokens redacted)
bb049db verified
Raw History Blame Contribute Delete
9.9 kB
#!/usr/bin/env python3
"""Visualize causal vs bidirectional ROI attention masks to illustrate the method."""
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import matplotlib.patches as mpatches
import numpy as np
def make_attention_matrix(seq_labels, bidir_range=None, bidir_mode='full'):
"""
Build attention matrix.
- Default: causal (lower triangular)
- bidir_range: (start, end) indices for bidirectional tokens
- bidir_mode:
'full' = bidir tokens attend to ALL tokens (past + future)
'mutual' = bidir tokens attend to each other + causal to rest
"""
n = len(seq_labels)
# Start with causal mask (lower triangular)
mask = np.tril(np.ones((n, n)))
if bidir_range is not None:
rs, re = bidir_range
if bidir_mode == 'full':
# ROI queries attend to full sequence
for i in range(rs, re):
mask[i, :] = 1.0
elif bidir_mode == 'mutual':
# ROI tokens only see each other bidirectionally
# Keep causal for non-ROI keys
for i in range(rs, re):
for j in range(rs, re):
mask[i, j] = 1.0 # ROI↔ROI bidir
return mask
def plot_mask(ax, mask, seq_labels, title, token_colors, bidir_cells=None):
"""Plot a single attention mask matrix with bidir cells highlighted."""
n = len(seq_labels)
# Create colored matrix
colored = np.ones((n, n, 3)) # white = masked
for i in range(n):
for j in range(n):
if mask[i, j] > 0:
colored[i, j] = token_colors[i]
ax.imshow(colored, aspect='equal', origin='upper')
# Mark bidir cells with a distinct pattern (small red dot)
if bidir_cells is not None:
for (i, j) in bidir_cells:
if mask[i, j] > 0:
ax.plot(j, i, 's', color='red', markersize=3, alpha=0.5)
# Grid lines
for i in range(n + 1):
ax.axhline(i - 0.5, color='gray', linewidth=0.5, alpha=0.4)
ax.axvline(i - 0.5, color='gray', linewidth=0.5, alpha=0.4)
ax.set_xticks(range(n))
ax.set_xticklabels(seq_labels, rotation=45, ha='right', fontsize=7)
ax.set_yticks(range(n))
ax.set_yticklabels(seq_labels, fontsize=7)
ax.set_xlabel('Key (attends to)', fontsize=9)
ax.set_ylabel('Query (token)', fontsize=9)
ax.set_title(title, fontsize=10, fontweight='bold', pad=10)
def main():
# === Inference sequence: [sys, visual(source), ROI crops, text query] ===
seq_labels = [
'sys₁', 'sys₂',
'v₁', 'v₂', 'v₃', 'v₄', 'v₅', 'v₆',
'roi₁', 'roi₂', 'roi₃', 'roi₄',
'q₁', 'q₂', 'q₃',
]
n = len(seq_labels)
color_map = {
'sys': np.array([0.7, 0.7, 0.7]),
'v': np.array([0.35, 0.55, 0.85]),
'roi': np.array([0.95, 0.55, 0.20]),
'q': np.array([0.40, 0.75, 0.40]),
'a': np.array([0.85, 0.35, 0.50]),
}
token_colors = []
for label in seq_labels:
if label.startswith('sys'): token_colors.append(color_map['sys'])
elif label.startswith('v'): token_colors.append(color_map['v'])
elif label.startswith('roi'): token_colors.append(color_map['roi'])
elif label.startswith('q'): token_colors.append(color_map['q'])
roi_start = seq_labels.index('roi₁')
roi_end = seq_labels.index('roi₄') + 1
# ===== Figure 1: 3-way comparison (inference) =====
fig, axes = plt.subplots(1, 3, figsize=(18, 6))
# (a) Standard causal
mask_causal = make_attention_matrix(seq_labels)
plot_mask(axes[0], mask_causal, seq_labels,
'(a) Standard Causal\n(baseline)', token_colors)
# (b) ROI↔ROI mutual bidir (ROI tokens see each other, causal to rest)
mask_mutual = make_attention_matrix(seq_labels,
bidir_range=(roi_start, roi_end),
bidir_mode='mutual')
# Collect bidir-only cells (above diagonal within ROI block)
bidir_cells_mutual = []
for i in range(roi_start, roi_end):
for j in range(roi_start, roi_end):
if j > i: # above diagonal = non-causal = the bidir addition
bidir_cells_mutual.append((i, j))
plot_mask(axes[1], mask_mutual, seq_labels,
'(b) ROI↔ROI Bidir\n(layers K→35)', token_colors,
bidir_cells=bidir_cells_mutual)
# Highlight ROI↔ROI block
rect1 = mpatches.FancyBboxPatch(
(roi_start - 0.5, roi_start - 0.5),
roi_end - roi_start, roi_end - roi_start,
linewidth=2.5, edgecolor='red', facecolor='none',
boxstyle='round,pad=0', linestyle='--'
)
axes[1].add_patch(rect1)
axes[1].annotate('mutual bidir\n(ROI↔ROI only)',
xy=(roi_end + 0.5, roi_start + 1), fontsize=8,
color='red', fontweight='bold')
# (c) ROI→ALL bidir (ROI sees everything including future)
mask_full = make_attention_matrix(seq_labels,
bidir_range=(roi_start, roi_end),
bidir_mode='full')
bidir_cells_full = []
for i in range(roi_start, roi_end):
for j in range(n):
if j > i:
bidir_cells_full.append((i, j))
plot_mask(axes[2], mask_full, seq_labels,
'(c) ROI→All Bidir\n(layers K→35)', token_colors,
bidir_cells=bidir_cells_full)
rect2 = mpatches.FancyBboxPatch(
(-0.5, roi_start - 0.5), n, roi_end - roi_start,
linewidth=2.5, edgecolor='red', facecolor='none',
boxstyle='round,pad=0', linestyle='--'
)
axes[2].add_patch(rect2)
axes[2].annotate('full bidir\n(ROI sees all)',
xy=(n - 1, roi_start + 1), fontsize=8,
color='red', fontweight='bold', ha='right')
# Legend
legend_patches = [
mpatches.Patch(color=color_map['sys'], label='System tokens'),
mpatches.Patch(color=color_map['v'], label='Visual tokens (source)'),
mpatches.Patch(color=color_map['roi'], label='ROI crop tokens'),
mpatches.Patch(color=color_map['q'], label='Text query tokens'),
mpatches.Patch(facecolor='white', edgecolor='gray', label='Masked (cannot attend)'),
plt.Line2D([0], [0], marker='s', color='red', markersize=6,
linestyle='', alpha=0.5, label='Bidir addition (non-causal)'),
]
fig.legend(handles=legend_patches, loc='lower center', ncol=6,
fontsize=8, bbox_to_anchor=(0.5, -0.02))
plt.suptitle('Selective Bidirectional Visual Attention for SD-RPN (Inference, Layers K→35)',
fontsize=13, fontweight='bold', y=1.02)
plt.tight_layout()
plt.savefig('/opt/tiger/thothvl_pretrain/visualize/attention_mask_comparison.png',
dpi=200, bbox_inches='tight', pad_inches=0.3)
print("Saved: visualize/attention_mask_comparison.png")
# ===== Figure 2: Training masks (layers 0→K-1 vs K→35) =====
train_labels = [
'sys₁', 'sys₂',
'v₁', 'v₂', 'v₃', 'v₄', 'v₅', 'v₆',
'q₁', 'q₂',
'a₁', 'a₂', 'a₃',
]
train_colors = []
for label in train_labels:
if label.startswith('sys'): train_colors.append(color_map['sys'])
elif label.startswith('v'): train_colors.append(color_map['v'])
elif label.startswith('q'): train_colors.append(color_map['q'])
elif label.startswith('a'): train_colors.append(color_map['a'])
vis_s = train_labels.index('v₁')
vis_e = train_labels.index('v₆') + 1
fig2, axes2 = plt.subplots(1, 2, figsize=(14, 6))
# Layers 0→K-1: standard causal
mask_train_early = make_attention_matrix(train_labels)
plot_mask(axes2[0], mask_train_early, train_labels,
'Layers 0→K-1: Standard Causal', train_colors)
# Layers K→35: visual↔visual mutual bidir
mask_train_late = make_attention_matrix(train_labels,
bidir_range=(vis_s, vis_e),
bidir_mode='mutual')
bidir_cells_train = []
for i in range(vis_s, vis_e):
for j in range(vis_s, vis_e):
if j > i:
bidir_cells_train.append((i, j))
plot_mask(axes2[1], mask_train_late, train_labels,
'Layers K→35: Visual↔Visual Bidir', train_colors,
bidir_cells=bidir_cells_train)
rect3 = mpatches.FancyBboxPatch(
(vis_s - 0.5, vis_s - 0.5),
vis_e - vis_s, vis_e - vis_s,
linewidth=2.5, edgecolor='red', facecolor='none',
boxstyle='round,pad=0', linestyle='--'
)
axes2[1].add_patch(rect3)
axes2[1].annotate('mutual bidir\n(vis↔vis)',
xy=(vis_e + 0.3, vis_s + 1), fontsize=8,
color='red', fontweight='bold')
legend_patches2 = [
mpatches.Patch(color=color_map['sys'], label='System'),
mpatches.Patch(color=color_map['v'], label='Visual'),
mpatches.Patch(color=color_map['q'], label='Query'),
mpatches.Patch(color=color_map['a'], label='Answer (target)'),
plt.Line2D([0], [0], marker='s', color='red', markersize=6,
linestyle='', alpha=0.5, label='Bidir addition'),
]
fig2.legend(handles=legend_patches2, loc='lower center', ncol=5,
fontsize=9, bbox_to_anchor=(0.5, -0.02))
plt.suptitle('Training: Bidirectional Visual Attention (K=24)',
fontsize=14, fontweight='bold', y=1.02)
plt.tight_layout()
plt.savefig('/opt/tiger/thothvl_pretrain/visualize/attention_mask_training.png',
dpi=200, bbox_inches='tight', pad_inches=0.3)
print("Saved: visualize/attention_mask_training.png")
if __name__ == '__main__':
main()