将WordPieceTokenizer加入Keras端到端模型时遇维度错误求助
问题解决:StartEndPacker 输入rank错误修复
问题根源
你定义的输入层shape=(1,)会让输入张量变成二维结构(batch_size, 1),每个样本被视为「包含1个字符串的数组」。经过WordPieceTokenizer处理后,输出的ragged张量会保留这个额外维度,最终变成三维结构(batch_size, 1, seq_len),直接触发了StartEndPacker要求输入为1或2维的检查错误。
修复方案
将输入层的shape参数改为(),让每个样本直接对应单个字符串,而非包裹在额外维度中:
tokenizer = keras_nlp.tokenizers.WordPieceTokenizer(...) start_packer = keras_nlp.layers.StartEndPacker(...) ... # 修正输入层:shape=() 表示每个样本是单个字符串 inputs = tf.keras.Input(shape=(), dtype="string") # 分词后得到二维ragged张量 (batch_size, seq_len) indices = tokenizer(inputs) # 此时输入符合StartEndPacker的维度要求 packed = start_packer(indices) outputs = model(packed) end_to_end_model = tf.keras.Model(inputs, outputs) # 预测时直接传入字符串列表即可 outputs = end_to_end_model.predict([ "JumbleOText", "nvagonnagesthis", "Written for a Human" ])
补充说明
- 调整输入层后,输入张量形状为
(batch_size,)(一维),分词器处理后输出二维ragged张量,完全匹配StartEndPacker的输入要求。 - 预测时无需给字符串添加额外维度,直接传入字符串列表即可完成推理。
内容的提问来源于stack exchange,提问作者FrozenKiwi
相关产品推荐
相关产品推荐

