Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -457,7 +457,7 @@ def __init__(self, config, dim, input_resolution, num_heads, drop_path_rate=0.0,
self.intermediate = MaskFormerSwinIntermediate(config, dim)
self.output = MaskFormerSwinOutput(config, dim)

def get_attn_mask(self, input_resolution, device=None):
def get_attn_mask(self, input_resolution, device=None, dtype=torch.float32):
"""Build the cyclic-shift attention mask for shifted-window MSA; returns None when shift_size is 0.

Each (h, w) position belongs to one of 9 cyclic-shift regions (3 along each axis), encoded
Expand All @@ -476,7 +476,7 @@ def get_attn_mask(self, input_resolution, device=None):
w_idx = torch.arange(width, device=device)
h_region = (h_idx >= height - self.window_size).long() + (h_idx >= height - self.shift_size).long()
w_region = (w_idx >= width - self.window_size).long() + (w_idx >= width - self.shift_size).long()
img_mask = (h_region[None, :, None, None] * 3 + w_region[None, None, :, None]).float()
img_mask = (h_region[None, :, None, None] * 3 + w_region[None, None, :, None]).to(dtype=dtype)
mask_windows = window_partition(img_mask, self.window_size)
mask_windows = mask_windows.view(-1, self.window_size * self.window_size)
attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
Expand Down Expand Up @@ -511,7 +511,9 @@ def forward(self, hidden_states, input_dimensions, output_attentions=False):
# partition windows
hidden_states_windows = window_partition(shifted_hidden_states, self.window_size)
hidden_states_windows = hidden_states_windows.view(-1, self.window_size * self.window_size, channels)
attn_mask = self.get_attn_mask((height_pad, width_pad), device=hidden_states_windows.device)
attn_mask = self.get_attn_mask(
(height_pad, width_pad), device=hidden_states_windows.device, dtype=hidden_states_windows.dtype
)

self_attention_outputs = self.attention(hidden_states_windows, attn_mask, output_attentions=output_attentions)

Expand Down
Loading