tools 1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- tools/__init__.py +14 -0
- tools/array.py +39 -0
- tools/data.py +90 -0
- tools/exp.py +165 -0
- tools/hydra/__init__.py +6 -0
- tools/modules.py +117 -0
- tools/numpy/__init__.py +8 -0
- tools/numpy/_f.py +60 -0
- tools/numpy/_utils.py +117 -0
- tools/os.py +67 -0
- tools/pandas/__init__.py +45 -0
- tools/plot/__init__.py +5 -0
- tools/plot/sklearn.py +28 -0
- tools/plot/utils.py +70 -0
- tools/random.py +35 -0
- tools/sklearn/__init__.py +2 -0
- tools/sklearn/metrics.py +82 -0
- tools/sklearn/model_selection.py +229 -0
- tools/sklearn/preprocessing.py +125 -0
- tools/stats/__init__.py +43 -0
- tools/tools.py +568 -0
- tools/torch/__init__.py +15 -0
- tools/torch/_pandas.py +12 -0
- tools/torch/data.py +81 -0
- tools/torch/estimator.py +65 -0
- tools/torch/federated_learning.py +376 -0
- tools/torch/layers.py +13 -0
- tools/torch/model.py +397 -0
- tools/torch/optim/__init__.py +2 -0
- tools/torch/optim/lr_scheduler.py +195 -0
- tools/torch/plot.py +13 -0
- tools/torch/utils.py +121 -0
- tools-1.0.dist-info/LICENSE.txt +21 -0
- tools-1.0.dist-info/METADATA +37 -0
- tools-1.0.dist-info/RECORD +37 -0
- tools-1.0.dist-info/WHEEL +5 -0
- tools-1.0.dist-info/top_level.txt +1 -0
tools/torch/estimator.py
ADDED
|
@@ -0,0 +1,65 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
import torch.nn as nn
|
|
3
|
+
import torch.optim as optim
|
|
4
|
+
import torch.utils.data as D
|
|
5
|
+
from sklearn.base import BaseEstimator
|
|
6
|
+
import sys
|
|
7
|
+
import time
|
|
8
|
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
9
|
+
|
|
10
|
+
class NNEstimator(BaseEstimator):
|
|
11
|
+
def __init__(self, model=None, criterion=None, val_dataset=None, epoch=10, lr=0.001, batch_size=32, patience=20, device=device, log_dir=None):
|
|
12
|
+
# All arguments need default values. (sklearn)
|
|
13
|
+
# All arguments must be fed into self.args = args
|
|
14
|
+
self.model = model
|
|
15
|
+
self.criterion = criterion
|
|
16
|
+
|
|
17
|
+
self.epoch = epoch
|
|
18
|
+
self.lr = lr
|
|
19
|
+
self.batch_size = batch_size
|
|
20
|
+
self.device = device
|
|
21
|
+
|
|
22
|
+
# Optional arguments
|
|
23
|
+
self.val_dataset = val_dataset
|
|
24
|
+
self.log_dir = log_dir
|
|
25
|
+
|
|
26
|
+
def fit(self, x, y):
|
|
27
|
+
'''
|
|
28
|
+
:param x: numpy.array
|
|
29
|
+
:param y: nump.array
|
|
30
|
+
'''
|
|
31
|
+
self.model.to(self.device)
|
|
32
|
+
op = optim.Adam(self.model.parameters(), lr=self.lr)
|
|
33
|
+
|
|
34
|
+
x = torch.as_tensor(x, torch.float32)
|
|
35
|
+
y = torch.as_tensor(y)
|
|
36
|
+
|
|
37
|
+
loader = D.DataLoader(dataset, batch_size=self.batch_size, shuffle=True)
|
|
38
|
+
for epoch in range(self.epoch):
|
|
39
|
+
for x, y in self.loader:
|
|
40
|
+
x, y = x.to(self.device), y.to(self.device)
|
|
41
|
+
y_hat = self.model(x)
|
|
42
|
+
loss = self.criterion(y, y_hat)
|
|
43
|
+
|
|
44
|
+
op.zero_grad()
|
|
45
|
+
loss.backward()
|
|
46
|
+
op.step()
|
|
47
|
+
|
|
48
|
+
if self.val_dataset is not None:
|
|
49
|
+
validation
|
|
50
|
+
else:
|
|
51
|
+
loss.item()
|
|
52
|
+
|
|
53
|
+
self.model.cpu()
|
|
54
|
+
|
|
55
|
+
# if self.log_dir is None:
|
|
56
|
+
# log_dir = open(path.EXP.join(time.strftime('%Y-%m-%d_%H-%M-%S')+'.txt'))
|
|
57
|
+
# default_stdout = sys.stdout
|
|
58
|
+
# sys.stdout = log_dir
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
# All parameters created here should be fed into self.args_ = args
|
|
62
|
+
|
|
63
|
+
def predict(self, x):
|
|
64
|
+
# return predicted value
|
|
65
|
+
return y
|
|
@@ -0,0 +1,376 @@
|
|
|
1
|
+
# %%
|
|
2
|
+
import torch
|
|
3
|
+
import torch.nn as nn
|
|
4
|
+
|
|
5
|
+
# TODO: migrate these functions to hideandseek
|
|
6
|
+
|
|
7
|
+
# %%
|
|
8
|
+
def transfer(model_source, model_target):
|
|
9
|
+
'''
|
|
10
|
+
model_source: Single nn.Model instance
|
|
11
|
+
model_target: Single nn.Model instance
|
|
12
|
+
|
|
13
|
+
for multiple source or targets, refer to aggregate() or distribute()
|
|
14
|
+
'''
|
|
15
|
+
for p_trg, p_src in zip(model_target.parameters(), model_source.parameters()):
|
|
16
|
+
# model_target's device
|
|
17
|
+
device = p_trg.device
|
|
18
|
+
|
|
19
|
+
# Clone data. Becareful not to copy hard link! (never copy pointers)
|
|
20
|
+
p_trg.data = torch.clone(p_src.data).to(device)
|
|
21
|
+
|
|
22
|
+
# # 1. Reset
|
|
23
|
+
# data = torch.zeros_like(p_trg).to(device)
|
|
24
|
+
# # p_trg.data[:] = 0
|
|
25
|
+
#
|
|
26
|
+
# # 2. Add weighted sum
|
|
27
|
+
# data += p_src.data.to(device)
|
|
28
|
+
# p_trg.data += p_src.data.to(device)
|
|
29
|
+
|
|
30
|
+
'''# TODO: p_trg.data = torch.clone(p_src.data)'''
|
|
31
|
+
|
|
32
|
+
def aggregate(model_source, model_target, weight = None):
|
|
33
|
+
'''
|
|
34
|
+
Aggregate parameters of the model
|
|
35
|
+
|
|
36
|
+
model_source: List of nn.Model instances
|
|
37
|
+
model_target: Single nn.Model instance
|
|
38
|
+
weights: default = None
|
|
39
|
+
List of numbers. Must match the number of model_source.
|
|
40
|
+
If None, then weight is set as 1/len(model_source)
|
|
41
|
+
'''
|
|
42
|
+
if weight is None:
|
|
43
|
+
weight = [1/len(model_source)] * len(model_source)
|
|
44
|
+
assert len(model_source) == len(weight), "length of model_source(%s) and weight(%s) does not match"%(len(model_source), len(weight))
|
|
45
|
+
|
|
46
|
+
for parameters in zip(model_target.parameters(), *[model.parameters() for model in model_source]):
|
|
47
|
+
p_trg = parameters[0]
|
|
48
|
+
p_src_tuple = parameters[1:]
|
|
49
|
+
# model_target's device
|
|
50
|
+
device = p_trg.device
|
|
51
|
+
|
|
52
|
+
# 1. Reset
|
|
53
|
+
data = torch.zeros_like(p_trg)
|
|
54
|
+
# p_trg.data[:] = 0
|
|
55
|
+
|
|
56
|
+
# 2. Add weighted sum
|
|
57
|
+
for p_src, w in zip(p_src_tuple, weight):
|
|
58
|
+
# p_trg.data += (w * p_src.data).to(device)
|
|
59
|
+
data += (w * p_src.data).to(device)
|
|
60
|
+
|
|
61
|
+
p_trg.data = data
|
|
62
|
+
|
|
63
|
+
def distribute(model_source, model_target):
|
|
64
|
+
'''
|
|
65
|
+
Distribute parameters of the model
|
|
66
|
+
|
|
67
|
+
model_source: Single nn.Model instance
|
|
68
|
+
model_target: List of nn.Model instances
|
|
69
|
+
'''
|
|
70
|
+
for parameters in zip(model_source.parameters(), *[model.parameters() for model in model_target]):
|
|
71
|
+
p_src = parameters[0]
|
|
72
|
+
p_trg_tuple = parameters[1:]
|
|
73
|
+
|
|
74
|
+
for p_trg in p_trg_tuple:
|
|
75
|
+
device = p_trg.device
|
|
76
|
+
# p_trg.data[:] = p_src.data.to(device)
|
|
77
|
+
p_trg.data = torch.clone(p_src.data).to(device)
|
|
78
|
+
|
|
79
|
+
def aggregate_grad(model_source, model_target):
|
|
80
|
+
'''
|
|
81
|
+
model_source: List of nn.Model instances
|
|
82
|
+
model_target: Single nn.Model instance
|
|
83
|
+
'''
|
|
84
|
+
for parameters in zip(model_target.parameters(), *[model.parameters() for model in model_source]):
|
|
85
|
+
p_trg = parameters[0]
|
|
86
|
+
p_src_tuple = parameters[1:]
|
|
87
|
+
# model_target's device
|
|
88
|
+
device = p_trg.device
|
|
89
|
+
|
|
90
|
+
for p_src in p_src_tuple:
|
|
91
|
+
p_trg.grad += p_src.grad.to(device)
|
|
92
|
+
|
|
93
|
+
def aggregate_all (model_source, model_target):
|
|
94
|
+
'''
|
|
95
|
+
Aggregate parameters and states of the model.
|
|
96
|
+
States
|
|
97
|
+
|
|
98
|
+
model_source: List of nn.Model instances
|
|
99
|
+
model_target: Single nn.Model instance
|
|
100
|
+
'''
|
|
101
|
+
device = next(model.parameters()).device
|
|
102
|
+
state_dict = {}
|
|
103
|
+
state_dict_list = []
|
|
104
|
+
for model in model_source:
|
|
105
|
+
for key, value in model.state_dict().items():
|
|
106
|
+
if 'weight' in key or 'bias' in key:
|
|
107
|
+
continue
|
|
108
|
+
if key not in state_dict:
|
|
109
|
+
state_dict[key] = value.clone().to(device)
|
|
110
|
+
else:
|
|
111
|
+
state_dict[key] += value.to(device)
|
|
112
|
+
|
|
113
|
+
n_model = len(model_source)
|
|
114
|
+
for key in state_dict:
|
|
115
|
+
state_dict[key] /= n_model
|
|
116
|
+
|
|
117
|
+
model_target.load_state_dict(state_dict)
|
|
118
|
+
|
|
119
|
+
def distribute_all(model_source, model_target):
|
|
120
|
+
'''
|
|
121
|
+
model_source: Single nn.Model instance
|
|
122
|
+
model_target: List of nn.Model instances
|
|
123
|
+
'''
|
|
124
|
+
for model_target_ in model_target:
|
|
125
|
+
model_target_.load_state_dict(model_source.state_dict())
|
|
126
|
+
|
|
127
|
+
def aggregate_state_dict(model_list, device):
|
|
128
|
+
'''
|
|
129
|
+
aggregate model values which are not parameters(weight&bias) but are updated
|
|
130
|
+
'''
|
|
131
|
+
state_dict = {}
|
|
132
|
+
state_dict_list = []
|
|
133
|
+
for model in model_list:
|
|
134
|
+
print(next(model.parameters()).device)
|
|
135
|
+
for key, value in model.state_dict().items():
|
|
136
|
+
if 'weight' in key or 'bias' in key:
|
|
137
|
+
continue
|
|
138
|
+
if key not in state_dict:
|
|
139
|
+
state_dict[key] = value.clone().to(device)
|
|
140
|
+
else:
|
|
141
|
+
state_dict[key] += value.to(device)
|
|
142
|
+
|
|
143
|
+
n_model = len(model_list)
|
|
144
|
+
for key in state_dict:
|
|
145
|
+
state_dict[key] /= n_model
|
|
146
|
+
|
|
147
|
+
return state_dict
|
|
148
|
+
|
|
149
|
+
def attentive_aggregate(model_source_list, model_target, step_size = 0.01, p = 2, optimizer = None):
|
|
150
|
+
'''
|
|
151
|
+
perform attentive aggregation from source -> target.
|
|
152
|
+
model_source_list: List of nn.Model instances
|
|
153
|
+
model_target: Single nn.Model instance
|
|
154
|
+
'''
|
|
155
|
+
for parameters in zip(model_target.parameters(), *[model.parameters() for model in model_source]):
|
|
156
|
+
p_trg = parameters[0]
|
|
157
|
+
p_src_tuple = parameters[1:]
|
|
158
|
+
# model_target's device
|
|
159
|
+
device = p_trg.device
|
|
160
|
+
|
|
161
|
+
# 1. s = | w_server - w_cleint |
|
|
162
|
+
delta = []
|
|
163
|
+
for p_src in p_src_tuple:
|
|
164
|
+
delta.append(p_trg - p_src.to(device))
|
|
165
|
+
delta = torch.as_tensor(delta).to(device)
|
|
166
|
+
assert len(delta) == len(p_src_tuple)
|
|
167
|
+
|
|
168
|
+
s_k = torch.norm(delta, p=p, dim = tuple(range(1, len(p_trg.shape)+1 ) ) ) # dim = (1,2) or (1,2,3) or (1,2,3,...)
|
|
169
|
+
assert len(s_k) == len(p_src_tuple)
|
|
170
|
+
|
|
171
|
+
# 2. attention = softmax(s, dim = client)
|
|
172
|
+
attention = torch.softmax(s_k, dim = 0)
|
|
173
|
+
|
|
174
|
+
# 3. store gradient
|
|
175
|
+
gradient = torch.sum(attention.expand(delta.shape[::-1]).T * delta, dim = 0)
|
|
176
|
+
p_trg.grad[:] = gradient
|
|
177
|
+
|
|
178
|
+
# 4. apply gradient
|
|
179
|
+
p_trg.data -= step_size * p_trg.grad
|
|
180
|
+
# (possible to use optimizer here?)
|
|
181
|
+
|
|
182
|
+
|
|
183
|
+
# # 1. s = | w_server - w_cleint |
|
|
184
|
+
# s_k = []
|
|
185
|
+
# for p_src in p_src_tuple:
|
|
186
|
+
# s_k.append(torch.norm(p_trg - p_src.to(device), p=p)) # matrix norm of the difference of layer weights
|
|
187
|
+
# s_k = torch.as_tensor(s_k).to(device)
|
|
188
|
+
# assert len(s_k) == len(p_src_tuple)
|
|
189
|
+
#
|
|
190
|
+
# # 2. attention = softmax(s, dim = client)
|
|
191
|
+
# attention = torch.softmax(s_k, dim = 0)
|
|
192
|
+
#
|
|
193
|
+
# # 3. store gradient
|
|
194
|
+
# p_trg
|
|
195
|
+
# gradient = []
|
|
196
|
+
# for a, p_src in zip(attention, p_src_tuple):
|
|
197
|
+
# gradient.append(a * (p_trg - p_src.to(device) )
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
# 4. apply gradient (either by adding or )
|
|
201
|
+
|
|
202
|
+
# %%
|
|
203
|
+
if __name__ == '__main__':
|
|
204
|
+
|
|
205
|
+
class A(nn.Module):
|
|
206
|
+
def __init__(self, i=1):
|
|
207
|
+
super().__init__()
|
|
208
|
+
self.x = nn.Linear(5,5)
|
|
209
|
+
self.x.weight.data = torch.full_like(self.x.weight, i)
|
|
210
|
+
self.x.bias.data = torch.full_like(self.x.bias, i)
|
|
211
|
+
|
|
212
|
+
def forward(self, x):
|
|
213
|
+
return
|
|
214
|
+
|
|
215
|
+
# %%
|
|
216
|
+
m1 = A(1)
|
|
217
|
+
m2 = A(2)
|
|
218
|
+
print(m1.x.weight)
|
|
219
|
+
print(m2.x.weight)
|
|
220
|
+
id(m1.x.weight)
|
|
221
|
+
id(m2.x.weight)
|
|
222
|
+
id(m1.x.weight.data)
|
|
223
|
+
id(m2.x.weight.data)
|
|
224
|
+
|
|
225
|
+
# Only the data is copied
|
|
226
|
+
m1.x.weight.data[:] = m2.x.weight.data
|
|
227
|
+
print(m1.x.weight)
|
|
228
|
+
print(m2.x.weight)
|
|
229
|
+
id(m1.x.weight)
|
|
230
|
+
id(m2.x.weight)
|
|
231
|
+
id(m1.x.weight.data)
|
|
232
|
+
id(m2.x.weight.data)
|
|
233
|
+
|
|
234
|
+
m1.x.weight.data += 1
|
|
235
|
+
print(m1.x.weight)
|
|
236
|
+
print(m2.x.weight)
|
|
237
|
+
id(m1.x.weight)
|
|
238
|
+
id(m2.x.weight)
|
|
239
|
+
id(m1.x.weight.data)
|
|
240
|
+
id(m2.x.weight.data)
|
|
241
|
+
|
|
242
|
+
# %%
|
|
243
|
+
# Change together. Linked
|
|
244
|
+
m1.x.weight.data = m2.x.weight.data
|
|
245
|
+
print(m1.x.weight)
|
|
246
|
+
print(m2.x.weight)
|
|
247
|
+
id(m1.x.weight)
|
|
248
|
+
id(m2.x.weight)
|
|
249
|
+
id(m1.x.weight.data)
|
|
250
|
+
id(m2.x.weight.data)
|
|
251
|
+
|
|
252
|
+
m1.x.weight.data += 1
|
|
253
|
+
print(m1.x.weight)
|
|
254
|
+
print(m2.x.weight)
|
|
255
|
+
id(m1.x.weight)
|
|
256
|
+
id(m2.x.weight)
|
|
257
|
+
id(m1.x.weight.data)
|
|
258
|
+
id(m2.x.weight.data)
|
|
259
|
+
|
|
260
|
+
# Not linked
|
|
261
|
+
m1.x.weight.data = torch.clone(m2.x.weight.data)
|
|
262
|
+
print(m1.x.weight)
|
|
263
|
+
print(m2.x.weight)
|
|
264
|
+
id(m1.x.weight)
|
|
265
|
+
id(m2.x.weight)
|
|
266
|
+
id(m1.x.weight.data)
|
|
267
|
+
id(m2.x.weight.data)
|
|
268
|
+
|
|
269
|
+
m1.x.weight.data += 1
|
|
270
|
+
print(m1.x.weight)
|
|
271
|
+
print(m2.x.weight)
|
|
272
|
+
id(m1.x.weight)
|
|
273
|
+
id(m2.x.weight)
|
|
274
|
+
id(m1.x.weight.data)
|
|
275
|
+
id(m2.x.weight.data)
|
|
276
|
+
|
|
277
|
+
# %%
|
|
278
|
+
# Not Linked
|
|
279
|
+
transfer(m1,m2)
|
|
280
|
+
print(m1.x.weight)
|
|
281
|
+
print(m2.x.weight)
|
|
282
|
+
id(m1.x.weight)
|
|
283
|
+
id(m2.x.weight)
|
|
284
|
+
id(m1.x.weight.data)
|
|
285
|
+
id(m2.x.weight.data)
|
|
286
|
+
|
|
287
|
+
m1.x.weight.data += 1
|
|
288
|
+
print(m1.x.weight)
|
|
289
|
+
print(m2.x.weight)
|
|
290
|
+
id(m1.x.weight)
|
|
291
|
+
id(m2.x.weight)
|
|
292
|
+
id(m1.x.weight.data)
|
|
293
|
+
id(m2.x.weight.data)
|
|
294
|
+
|
|
295
|
+
# %%
|
|
296
|
+
# Aggregate
|
|
297
|
+
m1=A(1)
|
|
298
|
+
m2=A(2)
|
|
299
|
+
m3=A(3)
|
|
300
|
+
aggregate([m1,m2], m3)
|
|
301
|
+
print(m1.x.weight)
|
|
302
|
+
print(m2.x.weight)
|
|
303
|
+
print(m3.x.weight)
|
|
304
|
+
|
|
305
|
+
# %%
|
|
306
|
+
# distribute
|
|
307
|
+
distribute(m3, [m1,m2])
|
|
308
|
+
print(m1.x.weight)
|
|
309
|
+
print(m2.x.weight)
|
|
310
|
+
print(m3.x.weight)
|
|
311
|
+
m1.x.weight.data += 1
|
|
312
|
+
m2.x.weight.data += 2
|
|
313
|
+
print(m1.x.weight)
|
|
314
|
+
print(m2.x.weight)
|
|
315
|
+
print(m3.x.weight)
|
|
316
|
+
|
|
317
|
+
# from model import DNN
|
|
318
|
+
# import torch
|
|
319
|
+
# import torch.nn as nn
|
|
320
|
+
# import torch.nn.functional as F
|
|
321
|
+
# import torch.optim as optim
|
|
322
|
+
# n_input = 5
|
|
323
|
+
# n_output = 10
|
|
324
|
+
# n_hidden_list = [10,20,30,20]
|
|
325
|
+
# n_model = 4
|
|
326
|
+
# activation = [nn.Sigmoid(), nn.ReLU(), nn.LeakyReLU(0.03), nn.Tanh(), nn.Identity()]
|
|
327
|
+
# model_list = [DNN(n_input, n_output, n_hidden_list, activation) for i in range(n_model)]
|
|
328
|
+
#
|
|
329
|
+
# model_source = model_list[1:]
|
|
330
|
+
# model_target = model_list[0]
|
|
331
|
+
# model_target.cuda()
|
|
332
|
+
#
|
|
333
|
+
# x = torch.randn(20,5)
|
|
334
|
+
# for model in model_list:
|
|
335
|
+
# device = next(model.parameters()).device
|
|
336
|
+
# y_hat = model(x.to(device))
|
|
337
|
+
# loss = torch.sum(y_hat)
|
|
338
|
+
# loss.backward()
|
|
339
|
+
#
|
|
340
|
+
# p=list(model_target.parameters())
|
|
341
|
+
# p[0].grad
|
|
342
|
+
# next(model_source[0].parameters()).grad
|
|
343
|
+
#
|
|
344
|
+
# g = torch.zeros(10,5)
|
|
345
|
+
# for model in model_list:
|
|
346
|
+
# g += next(model.parameters()).grad.cpu()
|
|
347
|
+
# print('cumsum of gradient:\n',g)
|
|
348
|
+
|
|
349
|
+
|
|
350
|
+
# %%
|
|
351
|
+
# class m(nn.Module):
|
|
352
|
+
# def __init__(self):
|
|
353
|
+
# super(m, self).__init__()
|
|
354
|
+
# self.x = nn.Linear(10,10)
|
|
355
|
+
# # %%
|
|
356
|
+
# model_source = [m().cuda() for i in range(5)]
|
|
357
|
+
# model_target = m()
|
|
358
|
+
# for i, model in enumerate(model_source):
|
|
359
|
+
# model.x.weight.data[:] = i
|
|
360
|
+
# # %%
|
|
361
|
+
# for model in model_source:
|
|
362
|
+
# print(model.x.weight)
|
|
363
|
+
# print(model_target.x.weight.data)
|
|
364
|
+
# # %%
|
|
365
|
+
# aggregation(model_source, model_target)
|
|
366
|
+
# # %%
|
|
367
|
+
# print(model_target.x.weight.data)
|
|
368
|
+
# # %%
|
|
369
|
+
# distribution(model_target, model_source)
|
|
370
|
+
# # %%
|
|
371
|
+
#
|
|
372
|
+
# y.x.weight.data
|
|
373
|
+
# weight = None
|
|
374
|
+
# for i in zip([t.parameters() for t in model_source]):
|
|
375
|
+
# print(i)
|
|
376
|
+
# torch.cuda.device_count()
|
tools/torch/layers.py
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
import torch.nn as nn
|
|
3
|
+
import torch.nn.functional as F
|
|
4
|
+
import torch.optim as optim
|
|
5
|
+
|
|
6
|
+
class attention(nn.Module):
|
|
7
|
+
def __init__(self, dim=1):
|
|
8
|
+
super(att_pool, self).__init__()
|
|
9
|
+
self.attention = nn.softmax()
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def forward(self, x):
|
|
13
|
+
return x*self.attention(x)
|