TensorFlow 1.13非Eager模式下遍历张量的实现方法咨询
在TensorFlow 1.13中非Eager模式下解决张量遍历/修改问题
嘿,我完全懂你在TensorFlow 1.13里碰到的这个麻烦——非Eager模式下确实没法直接像遍历普通Python列表那样去循环张量,毕竟TF1.x是计算图优先的框架,所有操作都得转化成计算图里的节点,而不是即时执行代码逻辑。不过不用急,咱们完全可以不用开启Eager模式,用TensorFlow原生的向量化操作来实现你要的功能。
先分析你的问题根源
你原来的代码里用Python的for循环去遍历batch_index,还直接调用assign方法修改张量值——这在TF1.x的计算图模式下是行不通的:一方面这些循环操作不会被加入到计算图中,另一方面静态计算图里的张量本身就不支持这种即时修改的方式。
解决方案:用向量化操作替代Python循环
我们可以用TensorFlow的tf.where、广播机制和掩码操作来实现你要的“把指定位置的logits设为inf”的逻辑,全程不用显式循环,完全符合TF1.x的计算图模式:
import tensorflow as tf import math with tf.Session(config=config) as sess: context = tf.placeholder(tf.int32, [args.batch_size, None]) mask = tf.placeholder(tf.int32, [args.batch_size, 2]) output = model.model(hparams=hparams, X=context) # 1. 提取每个样本的start和end位置 starts = mask[:, 0] ends = mask[:, 1] # 2. 生成序列维度的索引,用于判断每个位置是否在[start, end]范围内 seq_len = tf.shape(context)[1] seq_indices = tf.range(seq_len, dtype=tf.int32) # shape: [seq_len] # 3. 广播生成每个样本的序列位置掩码:标记哪些seq位置需要修改 # 把starts/ends扩展维度,和seq_indices广播匹配 in_range = tf.logical_and( seq_indices >= tf.expand_dims(starts, axis=1), seq_indices <= tf.expand_dims(ends, axis=1) ) # shape: [batch_size, seq_len] # 4. 生成词汇表维度的掩码:标记每个seq位置对应的vocab索引 vocab_size = tf.shape(output['logits'])[2] vocab_indices = context # shape: [batch_size, seq_len] vocab_range = tf.range(vocab_size, dtype=tf.int32) # shape: [vocab_size] # 广播匹配后,得到每个(batch, seq, vocab)位置是否需要设为inf的掩码 vocab_mask = tf.equal(tf.expand_dims(vocab_indices, axis=-1), vocab_range) # shape: [batch_size, seq_len, vocab_size] # 5. 合并两个掩码:只有同时满足在seq范围和对应vocab索引的位置才修改 final_mask = tf.logical_and(tf.expand_dims(in_range, axis=-1), vocab_mask) # 6. 更新logits:把final_mask为True的位置设为inf,其余保持原值 updated_logits = tf.where( final_mask, tf.fill(tf.shape(output['logits']), tf.constant(math.inf, dtype=tf.float32)), output['logits'] ) # 7. 用更新后的logits计算loss loss = tf.reduce_mean( tf.nn.sparse_softmax_cross_entropy_with_logits( labels=context[:, 1:], logits=updated_logits[:, :-1] ) )
关键说明
- 所有操作都是TensorFlow的图操作,会被加入到静态计算图中,完全适配TF1.x的非Eager模式。
- 用广播机制替代了Python循环,不仅符合TF的设计思路,还能利用GPU的并行计算能力,效率更高。
- 如果你的序列长度是固定的,还可以把
tf.shape换成静态形状,但用tf.shape能适配动态序列长度的场景,更灵活。
内容的提问来源于stack exchange,提问作者piccolo
相关产品推荐
相关产品推荐

