MCPcopy Create free account
hub / github.com/AtlasAnalyticsLab/AdaFisher / create_mask

Function create_mask

Image_Classification/src/models/swin.py:50–63  ·  view source on GitHub ↗
(window_size, displacement, upper_lower, left_right)

Source from the content-addressed store, hash-verified

48
49
50def create_mask(window_size, displacement, upper_lower, left_right):
51 mask = torch.zeros(window_size ** 2, window_size ** 2)
52
53 if upper_lower:
54 mask[-displacement * window_size:, :-displacement * window_size] = float('-inf')
55 mask[:-displacement * window_size, -displacement * window_size:] = float('-inf')
56
57 if left_right:
58 mask = rearrange(mask, '(h1 w1) (h2 w2) -> h1 w1 h2 w2', h1=window_size, h2=window_size)
59 mask[:, -displacement:, :, :-displacement] = float('-inf')
60 mask[:, :-displacement, :, -displacement:] = float('-inf')
61 mask = rearrange(mask, 'h1 w1 h2 w2 -> (h1 w1) (h2 w2)')
62
63 return mask
64
65
66def get_relative_distances(window_size):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected