-
Notifications
You must be signed in to change notification settings - Fork 14
Expand file tree
/
Copy pathmaxpool.asm
More file actions
143 lines (111 loc) · 3.19 KB
/
Copy pathmaxpool.asm
File metadata and controls
143 lines (111 loc) · 3.19 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
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
global maxpool, maxpool_backward
; rdx = input address
; rsi = output address
; rdi = input size
; rcx = channel size
; rbx = pool_argmax address
maxpool:
mov r8, rdi
shr r8, 1 ; divide by 2*2
mov r9, r8 ; r8 = j, r9 = i on grid
mov rax, r8 ; save for reseting
imul r10, rcx, 4
mov r11, r10
imul r11, rdi ; size of a layer of input
mov r12, r10
add r12, r11
mov r13, rcx ; channel cnt from rcx to zero
.loop:
vmovdqu32 zmm0, [rdx]
vmovdqu32 zmm1, [rdx + r10]
vmovdqu32 zmm2, [rdx + r11]
vmovdqu32 zmm3, [rdx + r12]
; save before overwriting zmm0
vmovdqa32 zmm16, zmm0
vmovdqa32 zmm17, zmm1
vmovdqa32 zmm18, zmm2
vmovdqa32 zmm19, zmm3
vmaxps zmm0, zmm0, zmm1
vmaxps zmm0, zmm0, zmm2
vmaxps zmm0, zmm0, zmm3
vmovdqu32 [rsi], zmm0 ; save
vcmpps k0, zmm0, zmm16, 0x0E ; from zmm0 -> index 0
vcmpps k1, zmm0, zmm17, 0x0E ; from zmm1 -> index 1
vcmpps k2, zmm0, zmm18, 0x0E ; from zmm2 -> index 2
vcmpps k3, zmm0, zmm19, 0x0E ; from zmm3 -> index 3
; i think i should not use that 0x0E
; the left zeroes are being ignored
kandw k4, k0, k1 ; k4 = k0 & k1
kandw k5, k2, k4 ; k5 = k0 & k1 & k2
knotw k6, k0
korw k1, k1, k6 ; k1 |= ~k0
knotw k6, k4
korw k2, k2, k6 ; k2 |= ~k4
knotw k6, k5
korw k3, k3, k6 ; k3 |= ~k5
kmovw [rbx], k0 ; store 16-bit mask k0 to memory at rbx
kmovw [rbx+2], k1
kmovw [rbx+4], k2
kmovw [rbx+6], k3
add rsi, 64 ; next output
add rbx, 8 ; next argmax
add rdx, 64 ; next input
sub r13, 16 ; next channel
jnz .loop
mov r13, rcx ; reset channel cnt
add rdx, r10 ; go to next column, stride 2
dec r8
jnz .loop
mov r8, rax
add rdx, r11
dec r9
jnz .loop
ret ; end
; rdx = grad_conv input address of maxpool
; rsi = grad_output address of maxpool
; rdi = input size
; rcx = channel size
; rbx = pool_argmax address
maxpool_backward:
mov r8, rdi
shr r8, 1 ; divide by 2*2
mov r9, r8 ; r8 = j, r9 = i on grid
mov rax, r8 ; save for reseting
imul r10, rcx, 4
mov r11, r10
imul r11, rdi ; size of a layer of input
mov r12, r10
add r12, r11
mov r13, rcx ; channel cnt from rcx to zero
.loop:
kmovw k1, [rbx]
kmovw k2, [rbx+2]
kmovw k3, [rbx+4]
kmovw k4, [rbx+6]
knotw k1, k1
knotw k2, k2
knotw k3, k3
knotw k4, k4
vmovdqu32 zmm0, [rsi]
vmovaps zmm1{k1}{z}, zmm0
vmovaps zmm2{k2}{z}, zmm0
vmovaps zmm3{k3}{z}, zmm0
vmovaps zmm4{k4}{z}, zmm0
vmovdqu32 [rdx], zmm1
vmovdqu32 [rdx + r10], zmm2
vmovdqu32 [rdx + r11], zmm3
vmovdqu32 [rdx + r12], zmm4
add rsi, 64 ; next output grade
add rbx, 8 ; next argmax
add rdx, 64 ; next input
sub r13, 16 ; next channel
jnz .loop
mov r13, rcx ; reset channel cnt
add rdx, r10 ; go to next column, stride 2
dec r8
jnz .loop
mov r8, rax
add rdx, r11
dec r9
jnz .loop
ret ; end