Download visualize/plot_attention_masks.py from TerryPei/GroundFlow: direct link, hf CLI and curl.
- Browser
- Download file 9.9 kB
-
https://huggingface.co/TerryPei/GroundFlow/resolve/main/visualize/plot_attention_masks.py
- Command line
-
hf download hf://TerryPei/GroundFlow/visualize/plot_attention_masks.py
-
curl -L -o plot_attention_masks.py https://huggingface.co/TerryPei/GroundFlow/resolve/main/visualize/plot_attention_masks.py
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() | |