Seq2Seq模型编码器同输入输出训练时Graph执行错误求助
问题描述
我正在训练一个Seq2Seq模型,需求是让编码器用相同的输入输出对训练,之后仅训练编码器并将其冻结,但移除解码器后无法实现。训练时抛出InvalidArgumentError: Graph execution error,已知这类错误多因指标输入与预期维度不匹配,但排查后仍无法定位问题。
Encoder2代码
class Encoder2(tf.keras.layers.Layer): def __init__(self, text_processor, units): super(Encoder2, self).__init__() self.text_processor = text_processor self.vocab_size = text_processor.vocabulary_size() self.units = units # The embedding layer converts tokens to vectors self.embedding = tf.keras.layers.Embedding(self.vocab_size, units, mask_zero=True) # The RNN layer processes those vectors sequentially. self.rnn = tf.keras.layers.Bidirectional( merge_mode='sum', layer=tf.keras.layers.GRU(units, # Return the sequence and state return_sequences=True, recurrent_initializer='glorot_uniform')) def call(self, x): shape_checker = ShapeChecker() shape_checker(x, 'batch s') # 2. The embedding layer looks up the embedding vector for each token. x = self.embedding(x) shape_checker(x, 'batch s units') # 3. The GRU processes the sequence of embeddings. x = self.rnn(x) shape_checker(x, 'batch s units') # 4. Returns the new sequence of embeddings. return x def convert_input(self, texts): texts = tf.convert_to_tensor(texts) if len(texts.shape) == 0: texts = tf.convert_to_tensor(texts)[tf.newaxis] context = self.text_processor(texts).to_tensor() context = self(context) return context
Decoder2代码
class Decoder2(tf.keras.layers.Layer): @classmethod def add_method(cls, fun): setattr(cls, fun.__name__, fun) return fun def __init__(self, text_processor, units): super(Decoder2, self).__init__() self.text_processor = text_processor self.word_to_id = tf.keras.layers.StringLookup( vocabulary=text_processor.get_vocabulary(), mask_token='', oov_token='[UNK]') self.id_to_word = tf.keras.layers.StringLookup( vocabulary=text_processor.get_vocabulary(), mask_token='', oov_token='[UNK]', invert=True) self.start_token = self.word_to_id('[START]') self.end_token = self.word_to_id('[END]') self.vocab_size = text_processor.vocabulary_size() self.units = units # 1. The embedding layer converts token IDs to vectors self.embedding = tf.keras.layers.Embedding(self.vocab_size, units, mask_zero=True) # 2. The RNN keeps track of what's been generated so far. self.rnn = tf.keras.layers.GRU(units, return_sequences=True, return_state=True, recurrent_initializer='glorot_uniform') # 3. The RNN output will be the query for the attention layer. self.attention = CrossAttention(units) # 4. This fully connected layer produces the logits for each # output token. self.output_layer = tf.keras.layers.Dense(self.vocab_size) @Decoder2.add_method def call(self, context, x, state=None, return_state=False): shape_checker = ShapeChecker() shape_checker(x, 'batch t') shape_checker(context, 'batch s units') # 1. Lookup the embeddings x = self.embedding(x) shape_checker(x, 'batch t units') # 2. Process the target sequence. x, state = self.rnn(x, initial_state=state) shape_checker(x, 'batch t units') # 3. Use the RNN output as the query for the attention over the context. x = self.attention(x, context) self.last_attention_weights = self.attention.last_attention_weights shape_checker(x, 'batch t units') shape_checker(self.last_attention_weights, 'batch t s') # Step 4. Generate logit predictions for the next token. logits = self.output_layer(x) shape_checker(logits, 'batch t target_vocab_size') if return_state: return logits, state else: return logits @Decoder2.add_method def tokens_to_text(self, tokens): words = self.id_to_word(tokens) result = tf.strings.reduce_join(words, axis=-1, separator=' ') result = tf.strings.regex_replace(result, '^ *\[START\] *', '') result = tf.strings.regex_replace(result, ' *\[END\] *$', '') return result @Decoder2.add_method def get_next_token(self, context, next_token, done, state, temperature = 0.0): logits, state = self( context, next_token, state = state, return_state=True) if temperature == 0.0: next_token = tf.argmax(logits, axis=-1) else: logits = logits[:, -1, :]/temperature next_token = tf.random.categorical(logits, num_samples=1) # If a sequence produces an `end_token`, set it `done` done = done | (next_token == self.end_token) # Once a sequence is done it only produces 0-padding. next_token = tf.where(done, tf.constant(0, dtype=tf.int64), next_token) return next_token, done, state @Decoder2.add_method def get_initial_state(self, context): batch_size = tf.shape(context)[0] start_tokens = tf.fill([batch_size, 1], self.start_token) done = tf.zeros([batch_size, 1], dtype=tf.bool) embedded = self.embedding(start_tokens) return start_tokens, done, self.rnn.get_initial_state(embedded)[0]
Translator2代码
class Translator2(tf.keras.Model): @classmethod def add_method(cls, fun): setattr(cls, fun.__name__, fun) return fun def __init__(self, units, context_text_processor, target_text_processor): super().__init__() # Build the encoder and decoder encoder = Encoder2(context_text_processor2, units) decoder = Decoder2(target_text_processor2, units) self.encoder = encoder self.decoder = decoder def call(self, inputs): context, x = inputs context = self.encoder(context) logits = self.decoder(context, x) # TODO(b/250038731): remove this try: # Delete the keras mask, so keras doesn't scale the loss+accuracy. del logits._keras_mask except AttributeError: pass return context
训练代码
model2 = Translator2(UNITS, context_text_processor, target_text_processor2) logits2 = model2((ex_context_tok2, ex_context_tok2)) model2.compile(optimizer='adam', loss=masked_loss, metrics=[masked_acc, masked_loss]) history2 = model2.fit( train_ds2.repeat(), epochs=200, steps_per_epoch = 100, validation_data=val_ds2, validation_steps = 20, callbacks=[tf.keras.callbacks.EarlyStopping(patience=5) ])
问题分析与解决方案
核心问题
- 输出维度不匹配:Translator2的
call方法返回的是编码器输出context(形状(batch, seq_len, units)),但损失函数masked_loss和指标masked_acc需要的是与目标token ID匹配的logits(形状(batch, seq_len, vocab_size)),维度完全不兼容,触发图执行错误。 - 模型结构与需求不符:你的目标是训练编码器做自编码(输入输出相同),但当前结构保留了解码器,且最终返回的不是可计算损失的输出。
- 变量名不一致:Translator2初始化时使用了未传入的
context_text_processor2,存在变量名错误。
解决方案一:修改模型适配自编码训练(仅训练编码器)
如果目标是让编码器学习输入的自表示,可直接去掉解码器,添加输出层映射回词汇表:
class Translator2(tf.keras.Model): def __init__(self, units, context_text_processor): super().__init__() self.encoder = Encoder2(context_text_processor, units) # 添加输出层,将编码器输出映射回词汇表维度 self.output_layer = tf.keras.layers.Dense(context_text_processor.vocabulary_size()) def call(self, inputs): # 自编码场景,输入仅为context context = self.encoder(inputs) logits = self.output_layer(context) try: del logits._keras_mask except AttributeError: pass return logits
修改训练代码:
model2 = Translator2(UNITS, context_text_processor) # 验证输出形状:(batch, seq_len, vocab_size) logits2 = model2(ex_context_tok2) model2.compile(optimizer='adam', loss=masked_loss, metrics=[masked_acc]) history2 = model2.fit( train_ds2.repeat(), epochs=200, steps_per_epoch = 100, validation_data=val_ds2, validation_steps = 20, callbacks=[tf.keras.callbacks.EarlyStopping(patience=5)] )
解决方案二:保留解码器但仅训练编码器
如果必须保留解码器结构,需调整输出并冻结编码器:
- 冻结编码器:
model2.encoder.trainable = False
- 修改Translator2的
call方法,返回解码器的logits:
def call(self, inputs): context, x = inputs context = self.encoder(context) logits = self.decoder(context, x) try: del logits._keras_mask except AttributeError: pass # 返回logits而非context return logits
- 确保训练数据的输入结构为
(input_tokens, target_tokens),且target_tokens与logits维度匹配。
额外检查点
确认masked_loss和masked_acc的实现是否正确处理序列维度:
def masked_loss(y_true, y_pred): loss_fn = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True, reduction='none') loss = loss_fn(y_true, y_pred) mask = tf.cast(y_true != 0, loss.dtype) loss *= mask return tf.reduce_sum(loss) / tf.reduce_sum(mask) def masked_acc(y_true, y_pred): y_pred = tf.argmax(y_pred, axis=-1) y_pred = tf.cast(y_pred, y_true.dtype) match = tf.cast(y_true == y_pred, tf.float32) mask = tf.cast(y_true != 0, tf.float32) return tf.reduce_sum(match) / tf.reduce_sum(mask)
同时检查train_ds2和val_ds2的输出结构,确保为(input, target)或((context, x), target),与模型输入和损失计算逻辑匹配。
内容的提问来源于stack exchange,提问作者shey
相关产品推荐
相关产品推荐

