| 78 | |
| 79 | |
| 80 | def window_partition(x, window_size): |
| 81 | B, H, W, C = x.shape |
| 82 | pad_h = (window_size - H % window_size) % window_size |
| 83 | pad_w = (window_size - W % window_size) % window_size |
| 84 | if pad_h > 0 or pad_w > 0: |
| 85 | x = F.pad(x, (0, 0, 0, pad_w, 0, pad_h)) |
| 86 | Hp, Wp = H + pad_h, W + pad_w |
| 87 | x = x.view(B, Hp // window_size, window_size, Wp // window_size, window_size, C) |
| 88 | windows = x.permute(0, 1, 3, 2, 4, 5).reshape(-1, window_size, window_size, C) |
| 89 | return windows, (Hp, Wp) |
| 90 | |
| 91 | |
| 92 | def window_unpartition(windows, window_size, pad_hw, hw): |