文章目录
设置冻结层有两种方式。
- (不推荐)是在搭建网络时,直接将某层的trainable设置为false,例如:
layers.Conv2D(filters1, (1, 1), trainable=False)(input_tensor)
- 在网络搭建完成时,遍历model.layer,然后将layer.trainable设置为False:
# 冻结网络倒数的3层
for layer in model.layers[:-3]:
print(layer.trainable)
layer.trainable = False
也可以根据layer.name来确定哪些层需要冻结,例如冻结最后一层和RNN层:
for layer in model.layers:
layerName=str(layer.name)
if layerName.startswith("RNN_") or layerName.startswith("Final_"):
layer.trainable=False
在网络搭建时,可以考虑最后一个分类层命名和分类数量关联,这样当分类数量方式变化时,model.load_weight(“weight.h5”,by_name=True)不会加载最后一层