-
Notifications
You must be signed in to change notification settings - Fork 2.2k
Expand file tree
/
Copy pathwindow_process.py
More file actions
63 lines (51 loc) · 1.97 KB
/
Copy pathwindow_process.py
File metadata and controls
63 lines (51 loc) · 1.97 KB
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
# --------------------------------------------------------
# Fused kernel for window process for SwinTransformer
# Copyright (c) 2022 Nvidia
# Licensed under The MIT License [see LICENSE for details]
# --------------------------------------------------------
import torch
import swin_window_process
class WindowProcess(torch.autograd.Function):
@staticmethod
def forward(ctx, input, B, H, W, C, shift_size, window_size):
output = swin_window_process.roll_and_window_partition_forward(input, B, H, W, C, shift_size, window_size)
ctx.B = B
ctx.H = H
ctx.W = W
ctx.C = C
ctx.shift_size = shift_size
ctx.window_size = window_size
return output
@staticmethod
def backward(ctx, grad_in):
B = ctx.B
H = ctx.H
W = ctx.W
C = ctx.C
shift_size = ctx.shift_size
window_size = ctx.window_size
grad_out = swin_window_process.roll_and_window_partition_backward(grad_in, B, H, W, C, shift_size, window_size)
return grad_out, None, None, None, None, None, None, None
class WindowProcessReverse(torch.autograd.Function):
@staticmethod
def forward(ctx, input, B, H, W, C, shift_size, window_size):
output = swin_window_process.window_merge_and_roll_forward(input, B, H, W, C, shift_size, window_size)
ctx.B = B
ctx.H = H
ctx.W = W
ctx.C = C
ctx.shift_size = shift_size
ctx.window_size = window_size
return output
@staticmethod
def backward(ctx, grad_in):
B = ctx.B
H = ctx.H
W = ctx.W
C = ctx.C
shift_size = ctx.shift_size
window_size = ctx.window_size
#grad_out = ctx.saved_tensors[0]
#grad_out = torch.zeros((B, H, W, C), dtype=dtype).cuda()
grad_out = swin_window_process.window_merge_and_roll_backward(grad_in, B, H, W, C, shift_size, window_size)
return grad_out, None, None, None, None, None, None, None