You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用tf.GradientTape进行预训练模型迁移学习不收敛问题排查

Why Your GradientTape Training Isn't Converging

The core issue here is how the loss is calculated in your custom training loop, paired with a missing parameter that affects layer behavior during training. Let's break it down:

Root Cause 1: Loss Reduction Mismatch

When using model.fit(), TensorFlow automatically averages the loss over each batch (using SUM_OVER_BATCH_SIZE reduction). However, when you use tf.GradientTape directly, the default reduction=AUTO for SparseCategoricalCrossentropy switches to summing the loss over the entire batch instead of averaging it.

This creates gradients scaled by your batch size (e.g., 32x larger for a batch size of 32). Pre-trained models like MobileNetV2 have weights tuned for gradual updates—these oversized gradients cause the model to diverge instead of learning smoothly.

When you set base_model.trainable=False, only the small final dense layer is trained. This layer can tolerate larger updates without collapsing, which is why that case works (though the loss is still higher than model.fit() because updates aren't optimized).

Root Cause 2: Missing training=True Parameter

In your original custom loop, you called model(data) without specifying training=True. This tells layers like batch normalization or dropout to use inference-mode behavior (e.g., not updating running stats), which breaks proper training dynamics and slows convergence.

How to Fix the Code

Here's the corrected version of your custom training loop that aligns with model.fit() behavior:

import tensorflow as tf
from tensorflow import keras
from tensorflow.keras import layers

# Build the model as before
base_model = keras.applications.MobileNetV2(input_shape=(96, 96, 3), include_top=False, pooling='avg')
x = base_model.outputs[0]
outputs = layers.Dense(10, activation=tf.nn.softmax)(x)
model = keras.Model(inputs=base_model.inputs, outputs=outputs)

# Fix 1: Explicitly set loss reduction to match model.fit()
loss_fn = keras.losses.SparseCategoricalCrossentropy(reduction=keras.losses.Reduction.SUM_OVER_BATCH_SIZE)
optimizer = keras.optimizers.Adam()
train_loss = keras.metrics.Mean()
train_acc = keras.metrics.SparseCategoricalAccuracy()

def train_step(data, labels):
    with tf.GradientTape() as gt:
        # Fix 2: Add training=True to enable training-mode layer behavior
        pred = model(data, training=True)
        loss = loss_fn(labels, pred)
    grads = gt.gradient(loss, model.trainable_variables)
    optimizer.apply_gradients(zip(grads, model.trainable_variables))
    train_loss(loss)
    train_acc(labels, pred)

# Reset metrics at the start of each epoch (critical for multi-epoch training)
train_loss.reset_states()
train_acc.reset_states()

for xs, ys in train_data:
    train_step(xs, ys)

print('train_loss = {:.3f}, train_acc = {:.3f}'.format(train_loss.result(), train_acc.result()))

Alternative: Manually Average the Loss

If you prefer to use per-sample loss calculations, you can compute the average manually instead of setting the reduction parameter:

loss_fn = keras.losses.SparseCategoricalCrossentropy(reduction=keras.losses.Reduction.NONE)

def train_step(data, labels):
    with tf.GradientTape() as gt:
        pred = model(data, training=True)
        per_sample_loss = loss_fn(labels, pred)
        loss = tf.reduce_mean(per_sample_loss)  # Average over the batch
    grads = gt.gradient(loss, model.trainable_variables)
    optimizer.apply_gradients(zip(grads, model.trainable_variables))
    train_loss(loss)
    train_acc(labels, pred)

Why This Works

By aligning the loss averaging with model.fit() and enabling training-mode behavior for layers, you ensure the model receives properly scaled gradients and uses correct training dynamics. This allows the pre-trained base model to learn gradually without destabilization, leading to convergence that matches your original model.fit() results.

内容的提问来源于stack exchange,提问作者Wilson

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.14 06:47:57