TensorFlow变量梯度不存在警告及Transformer训练异常求助
问题分析与解决方案
核心问题原因
你遇到的两个现象(输出固定、embedding层无梯度)本质是Token嵌入层的参数未参与损失计算,根源是PositionEmbedding层的用法错误:
- 默认情况下,
keras_nlp.PositionEmbedding接收的是序列位置索引张量(形状为(batch_size, sequence_length)的整数张量),而非Token嵌入张量。你直接传入token_embeddings(形状(batch_size, seq_len, embed_dim)),导致Token嵌入的输出没有被传递到后续TransformerEncoder层,自然无法参与损失反向传播,参数永远不会更新,所有输入的Token嵌入输出一致,最终模型输出固定。 - 训练准确率虚高大概率是因为模型输出固定,恰好匹配了分布极度不平衡的标签,或者损失计算逻辑存在偏差。
修正方案(兼容自定义嵌入场景)
你可以选择以下两种方式正确组合Token嵌入和位置嵌入,同时兼容图像Patch嵌入等自定义嵌入形式:
方案1:手动生成位置嵌入并与Token嵌入相加
import tensorflow as tf from tensorflow.keras.layers import Input, Embedding from keras_nlp.layers import PositionEmbedding, TransformerEncoder encoder_inputs = Input(shape=(encoder_inputs_size,), name="encoder_inputs") # 1. Token嵌入层(替换为Patch嵌入层时,只需修改这部分逻辑) token_embeddings = Embedding(input_dim=vocabulary_size, output_dim=embedding_dim)(encoder_inputs) # 2. 生成位置嵌入:先创建固定位置索引,再通过PositionEmbedding生成嵌入 positions = tf.range(start=0, limit=encoder_inputs_size, delta=1) position_embeddings = PositionEmbedding(sequence_length=encoder_inputs_size)(positions) # 3. 相加Token嵌入与位置嵌入 combined_embeddings = token_embeddings + position_embeddings # 4. 传入Transformer编码器 encoder_outputs = TransformerEncoder(intermediate_dim=intermediate_dim, num_heads=num_heads)(inputs=combined_embeddings) encoder = Model(encoder_inputs, encoder_outputs, name="encoder")
方案2:利用PositionEmbedding的merge_mode参数自动相加
import tensorflow as tf from tensorflow.keras.layers import Input, Embedding from keras_nlp.layers import PositionEmbedding, TransformerEncoder encoder_inputs = Input(shape=(encoder_inputs_size,), name="encoder_inputs") # 1. Token嵌入层(兼容自定义嵌入) token_embeddings = Embedding(input_dim=vocabulary_size, output_dim=embedding_dim)(encoder_inputs) # 2. 设置merge_mode='add',让PositionEmbedding自动生成位置嵌入并与输入的Token嵌入相加 combined_embeddings = PositionEmbedding(sequence_length=encoder_inputs_size, merge_mode='add')(token_embeddings) # 3. 传入Transformer编码器 encoder_outputs = TransformerEncoder(intermediate_dim=intermediate_dim, num_heads=num_heads)(inputs=combined_embeddings) encoder = Model(encoder_inputs, encoder_outputs, name="encoder")
关键说明
两种方案都能确保Token(或Patch)嵌入层的参数参与损失计算,梯度正常更新,解决输出固定的问题。如果后续替换为图像Patch嵌入,只需修改生成token_embeddings的逻辑,位置嵌入的组合方式完全通用。
内容的提问来源于stack exchange,提问作者Cerno
相关产品推荐
相关产品推荐

