如何对5D张量格式的3D图像实现类似tf.image.crop_and_resize的裁剪缩放?
实现3D图像的裁剪与缩放(对应5D张量)
好问题!tf.image.crop_and_resize确实是为2D图像设计的,没法直接处理你这种[batch, image_height, image_width, image_depth, 1]的5D 3D张量。不过咱们可以通过两种思路来实现类似的功能,下面给你详细拆解:
方案一:复用2D API分维度处理
核心思路是把3D问题拆成两次2D处理——先处理平面维度(H/W),再处理深度维度(D)。这种方法不用自己写复杂的插值逻辑,直接复用TensorFlow的成熟API。
步骤拆解:
处理H/W维度:
- 把5D张量
[B, H, W, D, 1]reshape成[B*D, H, W, 1],相当于把每个深度切片当成一张独立的2D图像。 - 提取boxes中的H/W部分:每个原始box的
[y1, x1, y2, x2],重复D次(对应每个深度切片),得到形状为[B*D, 4]的boxes。 - 调用
tf.image.crop_and_resize,指定crop_size=[crop_h, crop_w],得到[B*D, crop_h, crop_w, 1]的结果。 - 再reshape回
[B, crop_h, crop_w, D, 1]。
- 把5D张量
处理D维度:
- 把上面的结果reshape成
[B*crop_h*crop_w, D, 1, 1],把深度维度当成“高度”,后面两个维度凑成2D图像的格式。 - 提取boxes中的D部分:每个原始box的
[z1, 0, z2, 0](因为后面两个维度是1,x坐标范围固定为0),重复crop_h*crop_w次,得到形状为[B*crop_h*crop_w, 4]的boxes。 - 再次调用
tf.image.crop_and_resize,指定crop_size=[crop_d, 1],得到[B*crop_h*crop_w, crop_d, 1, 1]的结果。 - 最后reshape回目标形状
[B, crop_h, crop_w, crop_d, 1]。
- 把上面的结果reshape成
代码示例:
import tensorflow as tf def crop_and_resize_3d_v1(images, boxes, box_indices, crop_size): # images shape: [B, H, W, D, 1] # boxes shape: [num_boxes, 6] -> [y1, x1, z1, y2, x2, z2] (归一化到[0,1]) # box_indices shape: [num_boxes] # crop_size: [crop_h, crop_w, crop_d] B, H, W, D, _ = tf.shape(images) num_boxes = tf.shape(boxes)[0] # Step 1: 处理H/W维度 images_2d_hw = tf.reshape(images, [B*D, H, W, 1]) boxes_hw = tf.tile(tf.expand_dims(boxes[:, [0,1,3,4]], 1), [1, D, 1]) boxes_hw = tf.reshape(boxes_hw, [num_boxes*D, 4]) box_indices_hw = tf.tile(tf.expand_dims(box_indices, 1), [1, D]) box_indices_hw = tf.reshape(box_indices_hw, [num_boxes*D]) cropped_hw = tf.image.crop_and_resize( images_2d_hw, boxes_hw, box_indices_hw, crop_size[:2] ) cropped_hw = tf.reshape(cropped_hw, [B, crop_size[0], crop_size[1], D, 1]) # Step 2: 处理D维度 images_2d_d = tf.reshape(cropped_hw, [B*crop_size[0]*crop_size[1], D, 1, 1]) boxes_d = tf.tile(tf.expand_dims(tf.concat([boxes[:, [2]], tf.zeros([num_boxes,1]), boxes[:, [5]], tf.zeros([num_boxes,1])], axis=1), 1), [1, crop_size[0]*crop_size[1], 1]) boxes_d = tf.reshape(boxes_d, [num_boxes*crop_size[0]*crop_size[1], 4]) box_indices_d = tf.tile(tf.expand_dims(box_indices, 1), [1, crop_size[0]*crop_size[1]]) box_indices_d = tf.reshape(box_indices_d, [num_boxes*crop_size[0]*crop_size[1]]) cropped_d = tf.image.crop_and_resize( images_2d_d, boxes_d, box_indices_d, [crop_size[2], 1] ) final_cropped = tf.reshape(cropped_d, [B, crop_size[0], crop_size[1], crop_size[2], 1]) return final_cropped
方案二:自定义3D版crop_and_resize(直接实现三线性插值)
如果追求效率或者需要更灵活的控制,可以直接实现3D版本的逻辑,核心是计算每个输出voxel对应的输入空间坐标,然后用三线性插值获取值。
核心步骤:
- 归一化boxes:确保输入的boxes是归一化到
[0,1]范围的(和tf.image.crop_and_resize一致)。 - 生成输出网格:为每个裁剪后的3D区域生成voxel坐标网格。
- 映射到输入空间:把输出网格的坐标转换为输入图像中的对应坐标。
- 三线性插值:根据输入坐标周围的8个voxel值计算当前输出voxel的结果。
代码示例:
import tensorflow as tf def crop_and_resize_3d_v2(images, boxes, box_indices, crop_size): # images shape: [B, H, W, D, 1] # boxes shape: [num_boxes, 6] -> [y1, x1, z1, y2, x2, z2] (归一化到[0,1]) # box_indices shape: [num_boxes] # crop_size: [crop_h, crop_w, crop_d] B, H, W, D, _ = tf.shape(images) num_boxes = tf.shape(boxes)[0] crop_h, crop_w, crop_d = crop_size # 转换为绝对坐标 boxes_abs = tf.stack([ boxes[:,0] * tf.cast(H, tf.float32), boxes[:,1] * tf.cast(W, tf.float32), boxes[:,2] * tf.cast(D, tf.float32), boxes[:,3] * tf.cast(H, tf.float32), boxes[:,4] * tf.cast(W, tf.float32), boxes[:,5] * tf.cast(D, tf.float32), ], axis=1) # 生成输出网格坐标 grid_y, grid_x, grid_z = tf.meshgrid( tf.linspace(0.0, 1.0, crop_h), tf.linspace(0.0, 1.0, crop_w), tf.linspace(0.0, 1.0, crop_d), indexing='ij' ) grid = tf.stack([grid_y, grid_x, grid_z], axis=-1) grid = tf.tile(tf.expand_dims(grid, 0), [num_boxes, 1, 1, 1, 1]) # 映射到输入空间坐标 box_start = boxes_abs[:, :3] box_size = boxes_abs[:, 3:] - box_start input_coords = box_start[:, tf.newaxis, tf.newaxis, tf.newaxis, :] + grid * box_size[:, tf.newaxis, tf.newaxis, tf.newaxis, :] # 获取对应batch的图像 batch_images = tf.gather(images, box_indices) # 拆分坐标为整数和小数部分 y, x, z = tf.split(input_coords, 3, axis=-1) y0, y1 = tf.floor(y), tf.floor(y)+1 x0, x1 = tf.floor(x), tf.floor(x)+1 z0, z1 = tf.floor(z), tf.floor(z)+1 # 裁剪坐标范围 y0 = tf.clip_by_value(y0, 0, tf.cast(H-1, tf.float32)) y1 = tf.clip_by_value(y1, 0, tf.cast(H-1, tf.float32)) x0 = tf.clip_by_value(x0, 0, tf.cast(W-1, tf.float32)) x1 = tf.clip_by_value(x1, 0, tf.cast(W-1, tf.float32)) z0 = tf.clip_by_value(z0, 0, tf.cast(D-1, tf.float32)) z1 = tf.clip_by_value(z1, 0, tf.cast(D-1, tf.float32)) # 转换为整数索引 y0_i, y1_i = tf.cast(y0, tf.int32), tf.cast(y1, tf.int32) x0_i, x1_i = tf.cast(x0, tf.int32), tf.cast(x1, tf.int32) z0_i, z1_i = tf.cast(z0, tf.int32), tf.cast(z1, tf.int32) # 获取8个邻域点的值 batch_idx = tf.tile(tf.range(num_boxes)[:,tf.newaxis,tf.newaxis,tf.newaxis], [1,crop_h,crop_w,crop_d]) v000 = tf.gather_nd(batch_images, tf.concat([batch_idx, y0_i, x0_i, z0_i], axis=-1)) v001 = tf.gather_nd(batch_images, tf.concat([batch_idx, y0_i, x0_i, z1_i], axis=-1)) v010 = tf.gather_nd(batch_images, tf.concat([batch_idx, y0_i, x1_i, z0_i], axis=-1)) v011 = tf.gather_nd(batch_images, tf.concat([batch_idx, y0_i, x1_i, z1_i], axis=-1)) v100 = tf.gather_nd(batch_images, tf.concat([batch_idx, y1_i, x0_i, z0_i], axis=-1)) v101 = tf.gather_nd(batch_images, tf.concat([batch_idx, y1_i, x0_i, z1_i], axis=-1)) v110 = tf.gather_nd(batch_images, tf.concat([batch_idx, y1_i, x1_i, z0_i], axis=-1)) v111 = tf.gather_nd(batch_images, tf.concat([batch_idx, y1_i, x1_i, z1_i], axis=-1)) # 计算插值权重 dy, dx, dz = y-y0, x-x0, z-z0 w000 = (1-dy)*(1-dx)*(1-dz) w001 = (1-dy)*(1-dx)*dz w010 = (1-dy)*dx*(1-dz) w011 = (1-dy)*dx*dz w100 = dy*(1-dx)*(1-dz) w101 = dy*(1-dx)*dz w110 = dy*dx*(1-dz) w111 = dy*dx*dz # 加权求和得到结果 interpolated = w000*v000 + w001*v001 + w010*v010 + w011*v011 + w100*v100 + w101*v101 + w110*v110 + w111*v111 interpolated = tf.reshape(interpolated, [num_boxes, crop_h, crop_w, crop_d, 1]) return interpolated
方案选择建议
- 如果你的场景对代码简洁性要求高,不想写复杂的插值逻辑,方案一更适合,直接复用TensorFlow优化过的2D crop_and_resize,虽然多了几次reshape,但大部分场景下效率足够。
- 如果追求更高的性能,或者需要自定义插值逻辑(比如更换为最近邻插值),方案二更灵活,直接对3D张量操作,减少了reshape带来的开销。
内容的提问来源于stack exchange,提问作者Devin
相关产品推荐
相关产品推荐

