File size: 9,895 Bytes
bb049db
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
#!/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()