TensorFlow Scan中忽略微小变化及统计超阈值变化元素数的实现
嘿,你的思路已经摸到门道啦!不过要同时实现「过滤小于阈值的小变化」和「统计超阈值变化的次数」,得给TensorFlow scan的传递状态加个计数项才行——毕竟每一步扫描都需要记住两个关键信息:上一次修正后的值,还有已经统计了多少次超阈值变化。我给你调整好代码,直接就能用:
第一步:改扫描函数,加入计数状态
原来的函数只返回了修正后的值,现在咱们让它返回一个元组,把计数也带上:
import tensorflow as tf def scan_fn(prev, curr): thrshd = 0.05 # 从之前的状态里取出上一次修正的值,还有累计的计数 prev_corrected_val, count = prev # 计算当前值和上一次修正值的差值 diff = curr - prev_corrected_val # 过滤掉小于阈值的变化:只有差值超过0.05的部分才保留 pos_change = tf.nn.relu(diff - thrshd) # 正向超过阈值的部分 neg_change = tf.nn.relu(-diff - thrshd) # 负向超过阈值的部分 # 计算当前修正后的值 new_corrected_val = prev_corrected_val + pos_change - neg_change # 判断这次变化是否超阈值:只要正向/负向变化不为0,就计数+1 is_over_thresh = tf.cast( tf.logical_or(tf.greater(pos_change, 0.), tf.greater(neg_change, 0.)), tf.int32 ) new_count = count + is_over_thresh # 返回新的状态:(当前修正值, 新计数) return (new_corrected_val, new_count)
第二步:初始化扫描状态
扫描的初始状态要对应输入的形状:
- 初始修正值直接用输入序列的第一个元素(第一个值没有前值,直接保留)
- 初始计数设为0(第一个元素不算变化)
第三步:执行scan并提取结果
假设你的输入是二维张量(比如batch里的多个序列),咱们需要调整下维度,因为TensorFlow的scan默认对第一个维度扫描:
# 补全你的输入张量(示例数据) a = tf.constant([[.1, .26, .3, .2, .15], [.07, .35, .24, .22, .1]]) # 初始化状态:(每个序列的第一个元素, 初始计数0) initial_state = (a[:, 0], tf.zeros(tf.shape(a)[0], dtype=tf.int32)) # 转置输入:把序列维度放到第一个位置(scan默认扫第一个维度) transposed_a = tf.transpose(a, perm=[1, 0]) # 执行扫描,得到修正后的值序列和计数序列 corrected_vals_transposed, counts_transposed = tf.scan( scan_fn, transposed_a, initializer=initial_state ) # 转置回来恢复原形状 corrected_vals = tf.transpose(corrected_vals_transposed, perm=[1, 0]) counts = tf.transpose(counts_transposed, perm=[1, 0]) # 查看结果 print("修正后的值:") print(corrected_vals.numpy()) print("\n每个序列的超阈值变化次数:") print(counts.numpy())
关键细节唠两句
- 状态传递:用元组传递多个状态是TensorFlow scan实现多跟踪目标的标准操作,必须这么做才能同时保留修正值和计数
- 阈值判断:用
tf.logical_or判断正向/负向变化是否超阈值,转成int32后就能直接累加计数 - 维度处理:因为scan默认扫第一个维度,所以二维输入要先转置,把序列维度放到前面,扫完再转回来
内容的提问来源于stack exchange,提问作者John
相关产品推荐
相关产品推荐

