如何通过tf.data.Dataset为序列数据追加EOS元素?修正元素相加问题
问题分析与解决方案
你遇到的核心问题是:在tf.data.Dataset.map中,seq是Tensor对象而非Python列表,+运算符在这里会被解析为元素级加法,而非列表拼接。这就是为什么你得到的是每个元素都加1(因为char2int['EOS']=1),而不是在序列末尾追加EOS标记。
下面提供两种可行的修正方案:
方案一:在Tensor层面用tf.concat实现拼接
修改map函数中的逻辑,使用TensorFlow的tf.concat函数来完成序列与EOS标记的拼接,需要先把EOS转换成和seq维度一致的Tensor:
import tensorflow as tf def _get_generator(list_of_text, char2int): def gen(): for text in list_of_text: yield [char2int[x] for x in text] # transform char to int return gen def get_dataset(list_of_text, char2int): gen = _get_generator(list_of_text, char2int) dataset = tf.data.Dataset.from_generator(gen, (tf.int32), tf.TensorShape([None])) # 修正这里:用tf.concat替代+,先把EOS转为一维Tensor eos_tensor = tf.constant([char2int['EOS']], dtype=tf.int32) dataset = dataset.map(lambda seq: tf.concat([seq, eos_tensor], axis=0)) data_iter = dataset.make_initializable_iterator() return dataset, data_iter char2int = {'EOS':1, 'a':2, 'b':3, 'c':4} list_of_text = ['aaa', 'abc'] # the sequence data with tf.Graph().as_default(): dataset, data_iter = get_dataset(list_of_text, char2int) with tf.Session() as sess: sess.run(data_iter.initializer) tt1 = sess.run(data_iter.get_next()) tt2 = sess.run(data_iter.get_next()) print(tt1) # 输出 [2 2 2 1] print(tt2) # 输出 [2 3 4 1]
方案二:在生成器中直接追加EOS(更简洁)
既然生成器输出的是Python列表,我们可以直接在生成阶段就把EOS追加到每个序列末尾,这样后续Dataset不需要额外处理:
import tensorflow as tf def _get_generator(list_of_text, char2int): def gen(): for text in list_of_text: # 生成时直接追加EOS yield [char2int[x] for x in text] + [char2int['EOS']] return gen def get_dataset(list_of_text, char2int): gen = _get_generator(list_of_text, char2int) dataset = tf.data.Dataset.from_generator(gen, (tf.int32), tf.TensorShape([None])) data_iter = dataset.make_initializable_iterator() return dataset, data_iter char2int = {'EOS':1, 'a':2, 'b':3, 'c':4} list_of_text = ['aaa', 'abc'] # the sequence data with tf.Graph().as_default(): dataset, data_iter = get_dataset(list_of_text, char2int) with tf.Session() as sess: sess.run(data_iter.initializer) tt1 = sess.run(data_iter.get_next()) tt2 = sess.run(data_iter.get_next()) print(tt1) # 输出 [2 2 2 1] print(tt2) # 输出 [2 3 4 1]
两种方案都能得到你预期的结果,方案二更直接,减少了Tensor层面的操作开销。
内容的提问来源于stack exchange,提问作者Wilson
相关产品推荐
相关产品推荐

