#!/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()