求TensorFlow等价于Numpy的3D模型分块与重组函数(端到端训练)
Got it, let's break this down. The core issue here is that TensorFlow's reshape rearranges elements based on the tensor's memory layout, while your custom Numpy slicing function splits the 3D volume into spatially contiguous blocks along each axis. These two approaches have completely different element orderings, which is why reshape ruins your chair model's structure.
Below are TensorFlow implementations that exactly match the logic of your Numpy slice and combine functions—all operations are differentiable, so they'll work seamlessly for end-to-end training:
TensorFlow Equivalent of slice
import tensorflow as tf def tf_slice(obj, size): input_shape = obj.shape[0] # Generate start/end indices for each block along all axes axis_indices = tf.range(0, input_shape + 1, size) # Remove the last index (since it's beyond the input shape) axis_indices = axis_indices[:-1] blocks = [] # Iterate through all block positions across x, y, z axes for i_start, i_end in zip(axis_indices[:-1], axis_indices[1:]): for j_start, j_end in zip(axis_indices[:-1], axis_indices[1:]): for k_start, k_end in zip(axis_indices[:-1], axis_indices[1:]): # Extract the spatially contiguous block block = obj[i_start:i_end, j_start:j_end, k_start:k_end] blocks.append(block) # Stack all blocks into a single tensor of shape (8, 32, 32, 32) return tf.stack(blocks, axis=0)
TensorFlow Equivalent of combine
def tf_combine(lst, target_shape, size): # Initialize empty target tensor with the same dtype as input blocks reconstructed = tf.zeros((target_shape, target_shape, target_shape), dtype=lst.dtype) axis_indices = tf.range(0, target_shape + 1, size) axis_indices = axis_indices[:-1] block_idx = 0 # Iterate through the same block positions as slicing to reconstruct for i_start, i_end in zip(axis_indices[:-1], axis_indices[1:]): for j_start, j_end in zip(axis_indices[:-1], axis_indices[1:]): for k_start, k_end in zip(axis_indices[:-1], axis_indices[1:]): # Create slice indices for the target position # We use tensor_scatter_nd_update since TF tensors are immutable reconstructed = tf.tensor_scatter_nd_update( tensor=reconstructed, indices=[[[i_start, i_end], [j_start, j_end], [k_start, k_end]]], updates=[lst[block_idx]] ) block_idx += 1 return reconstructed
Key Notes:
- Unlike Numpy, TensorFlow tensors are immutable, so we use
tf.tensor_scatter_nd_updateinstead of direct assignment—this operation is fully differentiable, making it compatible with end-to-end training. - The indexing logic matches your Numpy code exactly: we split the 64x64x64 volume into 2x2x2 blocks of 32x32x32, iterating through each block position in the same order.
- To verify correctness, you can convert both the Numpy and TensorFlow outputs to Numpy arrays and compare them—they should be identical.
内容的提问来源于stack exchange,提问作者Fahd_
相关产品推荐
相关产品推荐

