-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathpm_v2.py
More file actions
67 lines (55 loc) · 2.78 KB
/
Copy pathpm_v2.py
File metadata and controls
67 lines (55 loc) · 2.78 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
import numpy as np
import tensorflow as tf
import tensorflow_probability as tfp
from scipy.stats import entropy
from network import MLP_rep
from dataset import Dataset
from baseline import Stream
tf.keras.backend.set_floatx('float64')
class PMStream(Stream):
def __init__(self, batch, epoch, sample_num, l):
super().__init__(batch, epoch, sample_num)
self.l = l
def train(self, train_ds, model, epoch):
for x_train, y_train in train_ds:
with tf.GradientTape() as tape:
# sampling
logit = model(x_train)
for s in range(self.sample_num - 1):
logit = tf.math.add(logit, model(x_train))
logit /= self.sample_num
neg_log_likelyhood = sum(self.loss_fn(y_true=y_train, y_pred=logit))
kl = sum(model.losses)
uq_all = - tf.math.reduce_sum(tf.math.log(logit + 1e-6) * logit, axis=1)
# i don't know good doing
# miss_idx = tf.equal(tf.dtypes.cast(tf.argmax(logit, axis=1), tf.int64), tf.dtypes.cast(tf.argmax(y_train, axis=1), tf.int64))
miss_idx = tf.math.logical_not(tf.equal(tf.dtypes.cast(tf.argmax(logit, axis=1), tf.int64), tf.dtypes.cast(tf.argmax(y_train, axis=1), tf.int64)))
uq_loss = tf.math.reduce_mean(tf.boolean_mask(uq_all, mask=miss_idx))
if np.isnan(uq_loss):
loss = ((kl + neg_log_likelyhood) / len(x_train))
else:
loss = ((kl + neg_log_likelyhood) / len(x_train)) - self.l * uq_loss
gradients = tape.gradient(loss, model.trainable_variables)
self.optimizer.apply_gradients(zip(gradients, model.trainable_variables))
self.train_score(y_train, logit)
self.loss_score(loss)
def run(self):
train_ds, test_ds = self.dataset.mnist()
model = tf.keras.models.load_model('baseline')
test_acc = list()
def map_epoch(epoch):
self.train_score.reset_states()
self.test_score.reset_states()
self.loss_score.reset_states()
self.train(train_ds, model, epoch)
self.test(test_ds, model)
print('epoch : {0}, train acc : {1:.4f}, train loss : {2:.4f}, test acc : {3:.4f}'.format(epoch, self.train_score.result(), self.loss_score.result(), self.test_score.result()))
test_acc.append(self.test_score.result().numpy())
list(map(map_epoch, range(self.epoch)))
model.save('pm_v2_{0}'.format(int(l * 10)))
np.save('pm_v2_acc_{0}'.format(int(l * 10)), np.array(test_acc))
if __name__ == '__main__':
ls = [0.05, 0.1, 0.5]
for l in ls:
stream = PMStream(128, 25, 10, l)
stream.run()