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
startandsizeparameters - Combine these ranges into a full grid of indices covering the entire slice
- Reshape the grid into the 2D index format that
scatter_nd_addexpects (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'intf.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_startandslice_sizeexactly as you would withtf.slice()—this works for any valid slice of your 3D tensor.
内容的提问来源于stack exchange,提问作者Vincent Stimper
相关产品推荐
相关产品推荐

