跟随TensorFlow NMT教程时遇InvalidArgumentError:输入形状不匹配
问题
我在Jupyter Notebook中跟随TensorFlow的seq2seq NMT(带注意力机制)教程学习,运行以下代码时:
# Setup the loop variables. next_token, done, state = decoder.get_initial_state(ex_context) tokens = [] for n in range(10): # Run one step. next_token, done, state = decoder.get_next_token( ex_context, next_token, done, state, temperature=1.0) # Add the token to the output. tokens.append(next_token) # Stack all the tokens together. tokens = tf.concat(tokens, axis=-1) # (batch, t) # Convert the tokens back to a a string result = decoder.tokens_to_text(tokens) result[:3].numpy()
触发了InvalidArgumentError,错误栈如下:
--------------------------------------------------------------------------- InvalidArgumentError Traceback (most recent call last) Cell In[31], line 2 1 # Setup the loop variables. ----> 2 next_token, done, state = decoder.get_initial_state(ex_context) 3 tokens = [] 5 for n in range(10): 6 # Run one step. Cell In[28], line 8 6 embedded = self.embedding(start_tokens) 7 print(embedded) ----> 8 return start_tokens, done, self.rnn.get_initial_state(embedded)[0] File ~/Library/Python/3.11/lib/python/site-packages/keras/src/layers/rnn/rnn.py:309, in RNN.get_initial_state(self, batch_size) 307 get_initial_state_fn = getattr(self.cell, "get_initial_state", None) 308 if get_initial_state_fn: --> 309 init_state = get_initial_state_fn(batch_size=batch_size) 310 else: 311 return [ 312 ops.zeros((batch_size, d), dtype=self.cell.compute_dtype) 313 for d in self.state_size 314 ] File ~/Library/Python/3.11/lib/python/site-packages/keras/src/layers/rnn/gru.py:326, in GRUCell.get_initial_state(self, batch_size) 324 def get_initial_state(self, batch_size=None): 325 return [ --> 326 ops.zeros((batch_size, self.state_size), dtype=self.compute_dtype) 327 ] File ~/Library/Python/3.11/lib/python/site-packages/keras/src/ops/numpy.py:5968, in zeros(shape, dtype) 5957 @keras_export(["keras.ops.zeros", "keras.ops.numpy.zeros"]) 5958 def zeros(shape, dtype=None): 5959 """Return a new tensor of given shape and type, filled with zeros. 5960 5961 Args: (...) 5966 Tensor of zeros with the given shape and dtype. 5967 """ --> 5968 return backend.numpy.zeros(shape, dtype=dtype) File ~/Library/Python/3.11/lib/python/site-packages/keras/src/backend/tensorflow/numpy.py:619, in zeros(shape, dtype) --> 617 return tf.zeros(shape, dtype=dtype) File ~/Library/Python/3.11/lib/python/site-packages/tensorflow/python/util/traceback_utils.py:153, in filter_traceback.<locals>.error_handler(*args, **kwargs) 151 except Exception as e: 152 filtered_tb = _process_traceback_frames(e.__traceback__) --> 153 raise e.with_traceback(filtered_tb) from None 154 finally: 155 del filtered_tb File ~/Library/Python/3.11/lib/python/site-packages/tensorflow/python/framework/ops.py:5983, in raise_from_not_ok_status(e, name) 5981 def raise_from_not_ok_status(e, name) -> NoReturn: 5982 e.message += (" name: " + str(name if name is not None else "")) --> 5983 raise core._status_to_exception(e) from None InvalidArgumentError: {{function_node __wrapped__Pack_N_2_device_/job:localhost/replica:0/task:0/device:CPU:0}} Shapes of all inputs must match: values[0].shape = [64,1,256] != values[1].shape = [] [Op:Pack] name:
我严格按照教程步骤操作,请问该如何解决?
解决方法
问题核心
错误源于新版本Keras中RNN层的get_initial_state接口变更:旧版本允许直接传入输入张量推导初始状态,而新版本仅接受batch_size参数。你的代码把形状为[64,1,256]的embedded张量直接传给该方法,导致batch_size被错误解析为张量而非标量,进而在创建零初始状态时引发形状不匹配。
修复代码
修改Decoder类中的get_initial_state方法,从embedded中提取正确的batch_size后传入RNN的初始化方法:
def get_initial_state(self, context): # 保留原有逻辑 batch_size = tf.shape(context)[0] start_tokens = tf.fill([batch_size, 1], self.target_vocab_size) done = tf.zeros([batch_size, 1], dtype=tf.bool) embedded = self.embedding(start_tokens) # 修改这一行:提取标量batch_size传入 rnn_initial_state = self.rnn.get_initial_state(batch_size=tf.shape(embedded)[0])[0] return start_tokens, done, rnn_initial_state
额外检查点
- 确认你的TensorFlow版本:若使用Keras 3.x或TF 2.15+,接口与旧版教程差异较大,需对应调整
- 验证
embedded的形状为[batch_size, seq_len, embed_dim],确保提取的batch_size是标量张量
内容的提问来源于stack exchange,提问作者CauseYNot
相关产品推荐
相关产品推荐

