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

TensorFlow 1.12:如何按起止与尺寸自动生成索引更新三维张量切片

Automatic Slice Index Generation for TensorFlow 1.12 Slice Assignment

Nice work getting the basic scatter_nd_add workflow running! To replicate the convenience of tf.slice() where you just specify a start position and slice size (instead of manually listing every index), we can build a helper function to generate the required indices automatically. Here's a practical, TF 1.12-compatible approach:

Core Approach

Instead of hardcoding indices, we'll:

  • Generate coordinate ranges for each dimension using your start and size parameters
  • Combine these ranges into a full grid of indices covering the entire slice
  • Reshape the grid into the 2D index format that scatter_nd_add expects (each row is a 3D coordinate tuple)

Implementation Code

import tensorflow as tf

def generate_3d_slice_indices(start, size):
    # Create ranges for each dimension based on start and size
    dim0_range = tf.range(start[0], start[0] + size[0])
    dim1_range = tf.range(start[1], start[1] + size[1])
    dim2_range = tf.range(start[2], start[2] + size[2])
    
    # Build a meshgrid of indices (use indexing='ij' to match tensor dimension order)
    dim0, dim1, dim2 = tf.meshgrid(dim0_range, dim1_range, dim2_range, indexing='ij')
    
    # Reshape each dimension to 1D and stack into the required 2D index tensor
    indices = tf.stack([
        tf.reshape(dim0, [-1]),
        tf.reshape(dim1, [-1]),
        tf.reshape(dim2, [-1])
    ], axis=1)
    
    return indices

# Test the solution with your original tensor shape
with tf.Session() as sess:
    init_val = tf.Variable(tf.zeros((2, 3, 3)))
    
    # Define slice parameters just like tf.slice()
    slice_start = [0, 0, 0]
    slice_size = [1, 2, 2]  # Matches the manual indices you used initially
    
    # Generate indices automatically
    indices = generate_3d_slice_indices(slice_start, slice_size)
    
    # Run the update operation
    update_op = tf.scatter_nd_add(init_val, indices, tf.ones(tf.shape(indices)[0]))
    
    sess.run(tf.global_variables_initializer())
    result = sess.run(update_op)
    
    print("Updated Tensor:\n", result)

Key Details

  • indexing='ij' in tf.meshgrid: This ensures we generate indices in tensor dimension order (first → second → third) instead of the default Cartesian order, which would mix up the axes.
  • Dynamic Size Handling: Using tf.shape(indices)[0] automatically sets the number of ones for the update, so you don't have to calculate the slice element count manually.
  • Flexibility: Adjust slice_start and slice_size exactly as you would with tf.slice()—this works for any valid slice of your 3D tensor.

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 08:34:07