forked from keras-team/keras
-
Notifications
You must be signed in to change notification settings - Fork 0
/
Copy pathtest_dynamic_trainability.py
110 lines (93 loc) · 3.34 KB
/
test_dynamic_trainability.py
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
from __future__ import absolute_import
from __future__ import print_function
import pytest
from keras.models import Model, Sequential
from keras.layers import Dense, Input
def test_layer_trainability_switch():
# with constructor argument, in Sequential
model = Sequential()
model.add(Dense(2, trainable=False, input_dim=1))
assert model.trainable_weights == []
# by setting the `trainable` argument, in Sequential
model = Sequential()
layer = Dense(2, input_dim=1)
model.add(layer)
assert model.trainable_weights == layer.trainable_weights
layer.trainable = False
assert model.trainable_weights == []
# with constructor argument, in Model
x = Input(shape=(1,))
y = Dense(2, trainable=False)(x)
model = Model(x, y)
assert model.trainable_weights == []
# by setting the `trainable` argument, in Model
x = Input(shape=(1,))
layer = Dense(2)
y = layer(x)
model = Model(x, y)
assert model.trainable_weights == layer.trainable_weights
layer.trainable = False
assert model.trainable_weights == []
def test_model_trainability_switch():
# a non-trainable model has no trainable weights
x = Input(shape=(1,))
y = Dense(2)(x)
model = Model(x, y)
model.trainable = False
assert model.trainable_weights == []
# same for Sequential
model = Sequential()
model.add(Dense(2, input_dim=1))
model.trainable = False
assert model.trainable_weights == []
def test_nested_model_trainability():
# a Sequential inside a Model
inner_model = Sequential()
inner_model.add(Dense(2, input_dim=1))
x = Input(shape=(1,))
y = inner_model(x)
outer_model = Model(x, y)
assert outer_model.trainable_weights == inner_model.trainable_weights
inner_model.trainable = False
assert outer_model.trainable_weights == []
inner_model.trainable = True
inner_model.layers[-1].trainable = False
assert outer_model.trainable_weights == []
# a Sequential inside a Sequential
inner_model = Sequential()
inner_model.add(Dense(2, input_dim=1))
outer_model = Sequential()
outer_model.add(inner_model)
assert outer_model.trainable_weights == inner_model.trainable_weights
inner_model.trainable = False
assert outer_model.trainable_weights == []
inner_model.trainable = True
inner_model.layers[-1].trainable = False
assert outer_model.trainable_weights == []
# a Model inside a Model
x = Input(shape=(1,))
y = Dense(2)(x)
inner_model = Model(x, y)
x = Input(shape=(1,))
y = inner_model(x)
outer_model = Model(x, y)
assert outer_model.trainable_weights == inner_model.trainable_weights
inner_model.trainable = False
assert outer_model.trainable_weights == []
inner_model.trainable = True
inner_model.layers[-1].trainable = False
assert outer_model.trainable_weights == []
# a Model inside a Sequential
x = Input(shape=(1,))
y = Dense(2)(x)
inner_model = Model(x, y)
outer_model = Sequential()
outer_model.add(inner_model)
assert outer_model.trainable_weights == inner_model.trainable_weights
inner_model.trainable = False
assert outer_model.trainable_weights == []
inner_model.trainable = True
inner_model.layers[-1].trainable = False
assert outer_model.trainable_weights == []
if __name__ == '__main__':
pytest.main([__file__])