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

TensorFlow批量图像操作:为整个批次应用相同偏移

3D医学影像全局统一随机偏移实现(兼容tf.data.map)

需求说明

这和TensorFlow批量图像操作的问题很相似,但区别在于处理3D医学影像时,需要对整个数据集的所有图像应用完全相同的随机偏移,而非每张图像单独生成偏移;同时实现必须兼容tf.data.Dataset的.map方法。

修改后的实现代码

import tensorflow as tf

class GlobalShift3DImages(object):
    def __init__(self, keys=('image', 'mask'), fill_value=None, fill_mode="reflect", interpolation="bilinear",
                 seed=None, depth_factor=0.0, height_factor=0.0, width_factor=0.0):
        """
        对3D医学影像数据集应用全局统一的随机偏移
        Args:
            depth_factor: 深度方向偏移比例,可为单个浮点数或二元组,表示上下限
                          负值表示向上偏移,正值表示向下偏移,如(-0.1, 0.1)表示偏移范围±10%
            height_factor: 高度方向偏移比例,规则同上
            width_factor: 宽度方向偏移比例,规则同上
            fill_mode: 边界外填充模式,可选"constant"、"nearest"、"wrap"、"reflect",默认"reflect"
            interpolation: 插值模式,可选"nearest"、"bilinear"
            seed: 随机种子,保证偏移量可复现
            fill_value: fill_mode为"constant"时的填充值
            keys: 需要同步偏移的数据集键名,如('image', 'mask')
        """
        self.keys = keys
        self.depth_factor = self._parse_factor(depth_factor)
        self.height_factor = self._parse_factor(height_factor)
        self.width_factor = self._parse_factor(width_factor)
        self.interpolation = interpolation
        self.fill_mode = fill_mode
        self.fill_value = fill_value
        self.seed = seed
        # 预先生成全局偏移量(初始化时不生成,第一次调用时生成)
        self._global_offsets = None

    def _parse_factor(self, factor):
        # 解析偏移比例为上下限二元组
        if isinstance(factor, (tuple, list)) and len(factor) == 2:
            return tuple(factor)
        elif isinstance(factor, float):
            return (-abs(factor), abs(factor))
        else:
            raise ValueError("factor必须是浮点数或长度为2的元组/列表")

    def _generate_global_offsets(self, image_shape):
        # 根据输入图像尺寸生成全局统一的偏移量
        depth, height, width = image_shape[0], image_shape[1], image_shape[2]
        tf.random.set_seed(self.seed)
        # 计算各方向的偏移像素值
        offset_depth = tf.random.uniform(
            shape=[], minval=self.depth_factor[0]*depth, maxval=self.depth_factor[1]*depth, dtype=tf.float32
        )
        offset_height = tf.random.uniform(
            shape=[], minval=self.height_factor[0]*height, maxval=self.height_factor[1]*height, dtype=tf.float32
        )
        offset_width = tf.random.uniform(
            shape=[], minval=self.width_factor[0]*width, maxval=self.width_factor[1]*width, dtype=tf.float32
        )
        self._global_offsets = (offset_depth, offset_height, offset_width)

    def _apply_translation(self, image):
        # 对单张3D图像应用预先生成的全局偏移
        depth, height, width = image.shape[0], image.shape[1], image.shape[2]
        # 创建平移变换矩阵
        transform = tf.convert_to_tensor([
            [1, 0, 0, self._global_offsets[2]],  # x轴(width)偏移
            [0, 1, 0, self._global_offsets[1]],  # y轴(height)偏移
            [0, 0, 1, self._global_offsets[0]],  # z轴(depth)偏移
            [0, 0, 0, 1]
        ], dtype=tf.float32)
        # 应用仿射变换
        shifted_image = tf.keras.layers.AffineTransform(
            transform=transform, interpolation=self.interpolation,
            fill_mode=self.fill_mode, fill_value=self.fill_value
        )(image)
        return shifted_image

    def shift(self, image_features):
        # 第一次调用时生成全局偏移量
        if self._global_offsets is None:
            # 取第一个键对应的图像尺寸作为基准
            sample_image = image_features[self.keys[0]]
            self._generate_global_offsets(sample_image.shape[:3])
        
        # 对每个键对应的图像应用相同偏移
        for key in self.keys:
            image_features[key] = self._apply_translation(image_features[key])
        return image_features

# 使用示例
if __name__ == "__main__":
    # 创建3D样本数据集:10个(64, 128, 128, 1)的图像和mask
    dataset = tf.data.Dataset.from_tensor_slices({
        'image': tf.random.uniform(shape=(10, 64, 128, 128, 1)),
        'mask': tf.random.uniform(shape=(10, 64, 128, 128, 1))
    })

    # 初始化全局偏移器,设置三个方向的偏移范围±10%
    global_shift = GlobalShift3DImages(
        depth_factor=0.1, height_factor=0.1, width_factor=0.1,
        fill_mode="constant", fill_value=0.0, seed=42
    )

    # 应用到数据集(注意:若需要每个epoch重新生成偏移,需重新实例化或重置_offsets)
    shifted_dataset = dataset.map(lambda x: global_shift.shift(x))

    # 验证:取前两个样本,检查偏移是否一致
    it = iter(shifted_dataset)
    sample1 = next(it)
    sample2 = next(it)
    # 对比图像的非零区域位置(简单验证)
    print("样本1图像偏移后的非零区域位置:", tf.where(sample1['image'] > 0)[:5])
    print("样本2图像偏移后的非零区域位置:", tf.where(sample2['image'] > 0)[:5])

关键改动说明

  • 全局偏移量生成:首次调用shift方法时,基于样本图像尺寸生成一次固定的偏移量,后续所有图像复用该偏移,保证全局一致性。
  • 3D扩展:新增depth_factor参数支持深度方向的偏移,使用3D仿射变换实现完整的3D图像平移。
  • 兼容.map方法:保持shift方法接收单样本字典输入,符合tf.data.Dataset.map的调用规范。
  • 可复现性:支持设置随机种子,确保每次生成的偏移量一致。

内容的提问来源于stack exchange,提问作者Brian Mark Anderson

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 14:37:32