You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用更简洁方法实现TensorFlow张量切片?能否用tf.gather类函数?

Question

I want to replicate the following numpy tensor slicing operation in TensorFlow, and I'm wondering if there's a more concise way to do it—preferably using tf.gather or tf.gather_nd.

Numpy Code

import numpy as np
a = np.reshape(np.arange(60), (3,2,2,5))
idx = np.array([0, 1, 0])
N = np.shape(a)[0]
mask = a[np.arange(N),:,:,idx]

I've tried several approaches, and only the following method worked, but it feels a bit cumbersome:

My Working TensorFlow Code

import tensorflow as tf
import numpy as np
a = tf.cast(tf.constant(np.reshape(np.arange(60), (3,2,2,5))), tf.int32)
idx2 = tf.constant([0, 1, 0])
fn = lambda i: a[i][:,:,idx2[i]]
idx = tf.range(tf.shape(a)[0])
masks = tf.map_fn(fn, idx)
with tf.Session() as sess:
    print(sess.run(a))
    print(sess.run(tf.shape(masks)))
    print(sess.run(masks))

Is there a simpler implementation? Can I use tf.gather or tf.gather_nd to achieve this? Thanks a lot!


Answer

Absolutely! You can pull off this slicing much more cleanly with either tf.gather (the more concise option here) or tf.gather_nd. Let’s walk through both approaches:

Approach 1: Using tf.gather

The core idea is to align the idx tensor with the first dimension of a, then gather the right elements along the last axis. We just need to reshape idx so it broadcasts correctly with the middle dimensions:

import tensorflow as tf
import numpy as np

a = tf.cast(tf.constant(np.reshape(np.arange(60), (3,2,2,5))), tf.int32)
idx = tf.constant([0, 1, 0])

# Expand idx to (3,1,1) so it matches the middle dimensions of a for broadcasting
idx_expanded = tf.expand_dims(tf.expand_dims(idx, 1), 1)
# Gather elements along the last dimension, using batch_dims to align with the first axis
mask = tf.gather(a, idx_expanded, axis=3, batch_dims=1)
# Remove the extra singleton dimension added by gather
mask = tf.squeeze(mask, axis=3)

with tf.Session() as sess:
    print(sess.run(tf.shape(mask)))  # Output: [3 2 2]
    print(sess.run(mask))

Approach 2: Using tf.gather_nd

If you prefer using tf.gather_nd, you’ll need to build a full index tensor that points to every element you want to extract. This involves creating grids for each dimension and combining them:

import tensorflow as tf
import numpy as np

a = tf.cast(tf.constant(np.reshape(np.arange(60), (3,2,2,5))), tf.int32)
idx = tf.constant([0, 1, 0])

# Get shape dimensions for index construction
batch_size = tf.shape(a)[0]
dim1 = tf.shape(a)[1]
dim2 = tf.shape(a)[2]

# Create repeated indices for each dimension
batch_idx = tf.tile(tf.reshape(tf.range(batch_size), (-1,1,1)), (1, dim1, dim2))
dim1_idx = tf.tile(tf.reshape(tf.range(dim1), (1,-1,1)), (batch_size, 1, dim2))
dim2_idx = tf.tile(tf.reshape(tf.range(dim2), (1,1,-1)), (batch_size, dim1, 1))
dim3_idx = tf.tile(tf.reshape(idx, (-1,1,1)), (1, dim1, dim2))

# Stack indices into a (3,2,2,4) tensor where each entry is a full coordinate
indices = tf.stack([batch_idx, dim1_idx, dim2_idx, dim3_idx], axis=-1)
# Gather the target elements
mask = tf.gather_nd(a, indices)

with tf.Session() as sess:
    print(sess.run(tf.shape(mask)))  # Output: [3 2 2]
    print(sess.run(mask))

Why these are better than tf.map_fn

  • They’re more efficient: tf.map_fn acts like a loop under the hood, which can slow down operations on large tensors. These built-in ops are optimized for performance.
  • They’re more idiomatic: They follow TensorFlow’s tensor-first design pattern, making your code easier to read and maintain for other TensorFlow developers.

Both methods produce exactly the same result as your original numpy code and working TensorFlow implementation. The tf.gather approach is the clear winner here for its brevity.

内容的提问来源于stack exchange,提问作者ppzhu

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.13 07:41:02