TensorFlow自定义训练循环中Dense层作为首层无法训练问题
TensorFlow自定义训练循环异常:Dense层作首层无法训练,加Flatten层恢复正常
问题描述
使用已扁平化的MNIST数据集训练模型时,若直接将Dense层作为网络首层,模型无法正常收敛训练;但在已扁平化的数据上额外添加Flatten层后,模型训练恢复正常。
注:在已扁平化数据上添加Flatten层仅为验证Dense层作为首层的问题,若在非扁平化数据上使用Conv2D层作为首层,模型训练完全正常,问题集中在Dense层的首层使用场景。
环境版本
- TensorFlow 2.9.1
- Python 3.8.6
可正常训练的模型
class CustomModel(keras.Model): def __init__(self, num_classes, name = None): super().__init__(name = name) self._flatten = tf.keras.layers.Flatten() self._dense1 = tf.keras.layers.Dense(64) self._dense2 = tf.keras.layers.Dense(num_classes) @tf.function def call(self, X, training=False): X = self._flatten(X) X = tf.nn.relu(self._dense1(X)) return self._dense2(X)
无法训练的模型
class CustomModel(keras.Model): def __init__(self, num_classes, name = None): super().__init__(name = name) self._dense1 = tf.keras.layers.Dense(64) self._dense2 = tf.keras.layers.Dense(num_classes) @tf.function def call(self, X, training=False): X = tf.nn.relu(self._dense1(X)) return self._dense2(X)
数据集处理代码
import tensorflow_datasets as tfds (ds_train, ds_test), ds_info = tfds.load( "mnist", split = ["train", "test"], shuffle_files = True, as_supervised = True, with_info = True ) def normalize_img(image, label): return tf.cast(image, tf.float32) / 255.0, label def flatten_img(image, label): return tf.reshape(image, [-1, 28 * 28]), label AUTOTUNE = tf.data.experimental.AUTOTUNE BATCH_SIZE = 64 # 训练数据集处理 ds_train = ds_train.map(normalize_img, num_parallel_calls = AUTOTUNE) ds_train = ds_train.map(flatten_img, num_parallel_calls = AUTOTUNE) ds_train = ds_train.cache() ds_train = ds_train.shuffle(ds_info.splits["train"].num_examples) ds_train = ds_train.batch(BATCH_SIZE) ds_train = ds_train.prefetch(AUTOTUNE) # 测试数据集处理 ds_test = ds_test.map(normalize_img, num_parallel_calls = AUTOTUNE) ds_test = ds_test.map(flatten_img, num_parallel_calls = AUTOTUNE) ds_test = ds_test.batch(BATCH_SIZE) ds_test = ds_test.prefetch(AUTOTUNE)
自定义训练循环代码
model = CustomModel(10) num_epochs = 5 optimizer = keras.optimizers.Adam() loss_fn = keras.losses.SparseCategoricalCrossentropy(from_logits=True) acc_metric = keras.metrics.SparseCategoricalAccuracy() @tf.function def train_epoch(x, y): with tf.GradientTape() as tape: y_pred = model(x, training = True) loss = loss_fn(y, y_pred) # 计算梯度 gradients = tape.gradient(loss, model.trainable_weights) # 反向传播更新权重 optimizer.apply_gradients(zip(gradients, model.trainable_weights)) acc_metric.update_state(y, y_pred) return loss # 训练循环 for epoch in range(num_epochs): print(f"\nStart of Training Epoch {epoch + 1}") for batch_idx, (x_batch, y_batch) in tqdm(enumerate(ds_train), total=len(ds_train)): loss = train_epoch(x_batch, y_batch) print(f"Accuracy :- {acc_metric.result()}, Loss :- {loss}") acc_metric.reset_states()
训练效果对比
- 带Flatten层的模型:训练正常,准确率逐步攀升,损失持续下降,符合预期
- 无Flatten层的模型:训练异常,准确率维持在随机猜测水平,损失无明显下降趋势
内容的提问来源于stack exchange,提问作者Prithwiraj Mitra
相关产品推荐
相关产品推荐

