TensorFlow模型输入不规则序列对的实现问题
问题解答
1. 能否执行x[:, 0, :]索引而不报错?
不能直接这么写。原因如下:
- 输入是RaggedTensor,经过
tf.keras.layers.Embedding后返回的仍是RaggedTensor,形状为(batch_size, 2, None, 16)(ragged维度为倒数第二个)。 - 执行
tf.reduce_sum(x, axis=-2)后,ragged维度被消除,结果变成普通Tensor,形状为(batch_size, 2, 16)。 - 但Keras自动生成的
SlicingOpLambda层仍尝试用RaggedTensor的索引逻辑处理这个普通Tensor,导致类型错误。
2. TensorFlow 2 API最规范的实现方式
推荐两种简洁且符合规范的实现方式:
方式一:拆分输入为两个独立的RaggedTensor
这种方式逻辑清晰,贴合Keras输入设计习惯:
import tensorflow as tf # 定义两个独立的可变长度序列输入 input1 = tf.keras.Input(shape=(None,), ragged=True, dtype=tf.int32) input2 = tf.keras.Input(shape=(None,), ragged=True, dtype=tf.int32) # 共享嵌入层 embedder = tf.keras.layers.Embedding(input_dim=16, output_dim=16) v1 = tf.reduce_sum(embedder(input1), axis=1) v2 = tf.reduce_sum(embedder(input2), axis=1) # 计算点积 outputs = tf.reduce_sum(tf.multiply(v1, v2), axis=1) model = tf.keras.models.Model(inputs=[input1, input2], outputs=[outputs]) model.compile(loss=tf.keras.losses.BinaryCrossentropy(from_logits=True)) # 准备数据集 xs1 = tf.ragged.constant([[0,1,2], [2,0]]) xs2 = tf.ragged.constant([[3,4], [5]]) ys = tf.constant([0, 1]) dataset = tf.data.Dataset.from_tensor_slices(((xs1, xs2), ys)) model.fit(dataset)
方式二:保留单个输入,用tf.split拆分张量
如果必须保留单个输入结构,可用tf.split避免索引错误:
import tensorflow as tf inputs = tf.keras.Input(shape=(2, None), ragged=True, dtype=tf.int32) embedder = tf.keras.layers.Embedding(input_dim=16, output_dim=16) x = embedder(inputs) # 对每个序列求和,消除ragged维度 x = tf.reduce_sum(x, axis=-2) # 拆分第二个维度的两个序列 v1, v2 = tf.split(x, num_or_size_splits=2, axis=1) # 去除多余维度 v1 = tf.squeeze(v1, axis=1) v2 = tf.squeeze(v2, axis=1) outputs = tf.reduce_sum(tf.multiply(v1, v2), axis=1) model = tf.keras.models.Model(inputs=[inputs], outputs=[outputs]) model.compile(loss=tf.keras.losses.BinaryCrossentropy(from_logits=True)) xs = tf.ragged.constant([ [[0, 1, 2], [3, 4]], [[2, 0], [5]], ]) ys = tf.constant([0, 1]) dataset = tf.data.Dataset.from_tensor_slices((xs, ys)) model.fit(dataset)
两种方式都能正确处理可变长度序列对,规避原代码中的类型错误,符合TensorFlow 2/Keras的API规范。
内容的提问来源于stack exchange,提问作者Maxim Blumental
相关产品推荐
相关产品推荐

