在TensorFlow中如何将每行首尾非零元素间的元素设为1?
问题描述
给定由0和1组成的张量A:
A = [[1,0,1,0,1,1,0], [0,1,0,1,0,1,0], [0,0,0,1,0,0,1], [0,1,1,1,0,0,0]]
需求是:将每行中第一个非零元素和最后一个非零元素之间的所有元素填充为1,最终得到:
A = [[1,1,1,1,1,1,0], [0,1,1,1,1,1,0], [0,0,0,1,1,1,1], [0,1,1,1,0,0,0]]
之前尝试的代码报错,原因有两个:
- TensorFlow的张量是不可变对象,无法直接通过索引赋值修改
- 切片逻辑错误:代码是对全局的索引切片,而非按每行单独处理首尾非零元素
解决方案
核心步骤
- 对每行分别定位第一个和最后一个非零元素的列索引
- 生成掩码矩阵标记需要填充为1的位置
- 基于掩码和原张量生成最终结果
代码实现
import tensorflow as tf def fill_between_nonzeros(input_tensor): rows, cols = input_tensor.shape # 1. 获取每行第一个非零元素的列索引,全0行设为超出范围的cols first_nonzero = tf.argmax(input_tensor, axis=1, output_type=tf.int32) has_nonzero = tf.reduce_any(input_tensor == 1, axis=1) first_nonzero = tf.where(has_nonzero, first_nonzero, tf.constant(cols, dtype=tf.int32)) # 2. 获取每行最后一个非零元素的列索引,全0行设为-1 reversed_input = tf.reverse(input_tensor, axis=[1]) last_nonzero_rev = tf.argmax(reversed_input, axis=1, output_type=tf.int32) last_nonzero = cols - 1 - last_nonzero_rev last_nonzero = tf.where(has_nonzero, last_nonzero, tf.constant(-1, dtype=tf.int32)) # 3. 生成每行的列索引矩阵 col_indices = tf.tile(tf.range(cols)[tf.newaxis, :], [rows, 1]) # 4. 生成掩码:列索引在首尾非零之间的位置为True mask = tf.logical_and(col_indices >= first_nonzero[:, tf.newaxis], col_indices <= last_nonzero[:, tf.newaxis]) # 5. 替换掩码位置为1,其余保留原值 result = tf.where(mask, tf.ones_like(input_tensor), input_tensor) return result # 测试示例 A = tf.constant([[1,0,1,0,1,1,0], [0,1,0,1,0,1,0], [0,0,0,1,0,0,1], [0,1,1,1,0,0,0]]) output = fill_between_nonzeros(A) print(output.numpy())
关键说明
- 用
tf.argmax快速定位每行第一个1的位置,反转张量后再用argmax就能得到最后一个1的位置 - 处理全0行的边界情况,避免掩码生成错误
- 全程用TensorFlow的张量操作,不直接修改原张量,符合其不可变特性
内容的提问来源于stack exchange,提问作者cikod
相关产品推荐
相关产品推荐

