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

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)
    ])

问题分析与解决方案

核心问题

  1. 输出维度不匹配:Translator2的call方法返回的是编码器输出context(形状(batch, seq_len, units)),但损失函数masked_loss和指标masked_acc需要的是与目标token ID匹配的logits(形状(batch, seq_len, vocab_size)),维度完全不兼容,触发图执行错误。
  2. 模型结构与需求不符:你的目标是训练编码器做自编码(输入输出相同),但当前结构保留了解码器,且最终返回的不是可计算损失的输出。
  3. 变量名不一致: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)]
)

解决方案二:保留解码器但仅训练编码器

如果必须保留解码器结构,需调整输出并冻结编码器:

  1. 冻结编码器:
model2.encoder.trainable = False
  1. 修改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
  1. 确保训练数据的输入结构为(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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 10:15:34