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

如何在TensorFlow中基于动态索引实现张量切片

Solution for Dynamic Indexing with tf.gather_nd

Got it, let's solve this problem step by step. The key here is to construct the full 3D index tensor dynamically by pairing each batch index with the corresponding sequence indices from your idx tensor. Here's how you can do it efficiently without relying on tf.map_fn:

Step-by-Step Implementation

First, let's start with your sample tensors to verify the solution works:

import tensorflow as tf

# Your input tensor (shape: (2, 4, 4))
input_tensor = tf.constant([
    [[0.6036284, 0.0281072, 0.78739774, 0.79748493],
     [0.92121416, 0.31211454, 0.75201935, 0.49418229],
     [0.99500716, 0.35610485, 0.78246456, 0.32932794],
     [0.44941011, 0.33340591, 0.56897491, 0.16929366]],
    [[0.82108098, 0.50557786, 0.76569009, 0.04855939],
     [0.55340368, 0.11384677, 0.63739866, 0.09481387],
     [0.52711403, 0.5621863, 0.44211769, 0.85780412],
     [0.15423198, 0.80663997, 0.86868405, 0.48221472]]
])

# Your index tensor (shape: (2, 2))
idx = tf.constant([[2, 0], [2, 0]])

Now, construct the full index tensor:

# Get the batch size (first dimension of idx)
batch_size = tf.shape(idx)[0]

# Create batch indices: shape (2, 1) -> [[0], [1]]
batch_indices = tf.expand_dims(tf.range(batch_size), axis=1)

# Combine batch indices with sequence indices from idx
# Broadcasting will automatically expand batch_indices to (2, 2)
full_indices = tf.concat([batch_indices, tf.expand_dims(idx, axis=-1)], axis=-1)

# Now use tf.gather_nd to slice the input tensor
output = tf.gather_nd(input_tensor, full_indices)

Let's check the output shape and values:

print(output.shape)  # Should be (2, 2, 4)
print(output.numpy())

This will give you exactly the output you're expecting:

[[[0.99500716 0.35610485 0.78246456 0.32932794]
  [0.6036284  0.0281072  0.78739774 0.79748493]]

 [[0.52711403 0.5621863  0.44211769 0.85780412]
  [0.82108098 0.50557786 0.76569009 0.04855939]]]

How It Works

Let's break down the key parts:

  • Batch Indices: tf.range(batch_size) creates a tensor [0, 1] for your 2-batch example. We expand it to (2,1) so it can broadcast with the (2,2) idx tensor.
  • Full Index Construction: By concatenating batch_indices (expanded to match idx's shape via broadcasting) with idx (expanded to have a final dimension of 1), we get a (2,2,2) tensor. Each entry [batch_idx, seq_idx] tells tf.gather_nd exactly which slice to pick from the input.
  • Dynamic Compatibility: This approach works for any dynamic batch size or number of indices per batch. For example, if idx was (3,5) and input_tensor was (3,10,7), the output would automatically be (3,5,7).

Why This Is Better Than tf.map_fn

Using vectorized operations like tf.concat and broadcasting is far more efficient than tf.map_fn in TensorFlow. Map functions loop over each batch element, which can slow down execution, especially on large tensors. The vectorized approach leverages TensorFlow's optimized backend operations for better performance.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.29 06:45:26