如何在TensorFlow v1中基于张量对数据进行索引分组?
在TensorFlow v1中基于mod_labels拆分张量的解决方法
在TensorFlow v1中,无法直接用Python循环遍历张量(包括占位符)的元素,必须借助TensorFlow原生的图内操作完成筛选拆分。可以通过tf.equal生成掩码、tf.where获取索引、tf.gather提取对应元素的组合实现需求,具体代码如下:
示例代码
import tensorflow as tf # 定义张量(支持常量或占位符,此处以常量为例) mod_labels = tf.constant([0,1,1,0,0]) feats = tf.constant([[1,2,1], [3,2,6], [1,1,1], [9,8,4], [5,4,8]]) labels = tf.constant([1,53,12,89,54]) # 处理mod_labels=0的情况 mod0_mask = tf.equal(mod_labels, 0) mod0_indices = tf.where(mod0_mask) mod0_indices = tf.squeeze(mod0_indices, axis=1) # 将二维索引转为一维 mod0_feats = tf.gather(feats, mod0_indices) mod0_labels = tf.gather(labels, mod0_indices) # 处理mod_labels=1的情况 mod1_mask = tf.equal(mod_labels, 1) mod1_indices = tf.where(mod1_mask) mod1_indices = tf.squeeze(mod1_indices, axis=1) mod1_feats = tf.gather(feats, mod1_indices) mod1_labels = tf.gather(labels, mod1_indices) # 运行验证结果 with tf.Session() as sess: print("mod0_feats:\n", sess.run(mod0_feats)) print("mod0_labels:\n", sess.run(mod0_labels)) print("mod1_feats:\n", sess.run(mod1_feats)) print("mod1_labels:\n", sess.run(mod1_labels))
代码说明
tf.equal(mod_labels, 0):生成布尔掩码,标记mod_labels中值为0的位置tf.where(mod0_mask):获取掩码中True对应的索引,返回二维张量(每个元素是单个一维索引)tf.squeeze:将二维索引压缩为一维,匹配tf.gather的输入格式要求tf.gather(feats, mod0_indices):根据索引从feats中提取对应行,得到目标张量
如果你的mod_labels、feats、labels是占位符,上述代码完全适用——所有操作均为TensorFlow图内运算,不涉及Python层面的遍历操作。
内容的提问来源于stack exchange,提问作者saad_saeed
相关产品推荐
相关产品推荐

