如何用更简洁方法实现TensorFlow张量切片?能否用tf.gather类函数?
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!
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_fnacts 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

