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

如何对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。

步骤拆解:

  1. 处理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]。
  2. 处理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]。

代码示例:

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对应的输入空间坐标,然后用三线性插值获取值。

核心步骤:

  1. 归一化boxes:确保输入的boxes是归一化到[0,1]范围的(和tf.image.crop_and_resize一致)。
  2. 生成输出网格:为每个裁剪后的3D区域生成voxel坐标网格。
  3. 映射到输入空间:把输出网格的坐标转换为输入图像中的对应坐标。
  4. 三线性插值:根据输入坐标周围的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 09:29:49