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

如何在TensorFlow中实现图像随机裁剪粘贴自定义Keras层并解决张量赋值报错

解决TensorFlow自定义Cut-Paste增强层的赋值错误问题

问题根源

TensorFlow的EagerTensor是不可变对象,不支持直接索引赋值操作;同时原代码使用Python标准库的random模块,在TensorFlow图模式下会导致随机值固定(图构建阶段仅执行一次),无法实现每个样本的独立随机增强。

修正后的Cut-Paste实现

import tensorflow as tf
from tensorflow import keras
from tensorflow.keras import layers

class Cut_Paste(layers.Layer):
    def __init__(self, x_scale=10, y_scale=10, IMG_SIZE=(224,224), **kwargs):
        super().__init__(**kwargs)
        """
        定义裁剪区域的宽高范围
        x_scale和y_scale为图像宽高的百分比
        """        
        self.size_x, self.size_y = IMG_SIZE
        self.span_x = int(x_scale * self.size_x * 0.01)
        self.span_y = int(y_scale * self.size_y * 0.01)
        
    # 生成裁剪/粘贴区域的顶点坐标,支持批量样本
    def get_vertices(self, batch_size):
        fraction = 0.25
        # 计算中心区域的坐标范围
        x_min = int(self.size_x * 0.5 * (1 - fraction))
        x_max = int(self.size_x * 0.5 * (1 + fraction))
        # 为每个样本生成独立的x顶点
        vert_x = tf.random.uniform(
            shape=[batch_size], 
            minval=x_min, 
            maxval=x_max, 
            dtype=tf.int32
        )
        
        y_min = int(self.size_y * 0.5 * (1 - fraction))
        y_max = int(self.size_y * 0.5 * (1 + fraction))
        # 为每个样本生成独立的y顶点
        vert_y = tf.random.uniform(
            shape=[batch_size], 
            minval=y_min, 
            maxval=y_max, 
            dtype=tf.int32
        )
        
        start_x = vert_x - self.span_x // 2
        start_y = vert_y - self.span_y // 2
        end_x = vert_x + self.span_x // 2
        end_y = vert_y + self.span_y // 2
        
        return start_x, start_y, end_x, end_y
        
    def call(self, image, training=None):
        # 仅在训练阶段应用增强
        if training is None:
            training = tf.keras.backend.learning_phase()
        if not training:
            return image
            
        batch_size = tf.shape(image)[0]
        
        # 获取裁剪区域坐标
        cut_start_x, cut_start_y, cut_end_x, cut_end_y = self.get_vertices(batch_size)
        
        # 为每个样本裁剪子图
        def crop_single_image(args):
            img, s_x, s_y, e_x, e_y = args
            return tf.slice(img, [s_x, s_y, 0], [e_x - s_x, e_y - s_y, -1])
            
        sub_images = tf.map_fn(
            crop_single_image,
            (image, cut_start_x, cut_start_y, cut_end_x, cut_end_y),
            dtype=image.dtype
        )
        
        # 获取粘贴区域坐标
        paste_start_x, paste_start_y, paste_end_x, paste_end_y = self.get_vertices(batch_size)
        
        # 将子图粘贴到每个样本的目标位置
        def paste_single_image(args):
            img, sub_img, p_sx, p_sy, p_ex, p_ey = args
            height, width = tf.shape(img)[0], tf.shape(img)[1]
            
            # 生成粘贴区域的掩码
            x_coords = tf.range(height)[:, tf.newaxis]
            y_coords = tf.range(width)[tf.newaxis, :]
            mask_x = tf.logical_and(x_coords >= p_sx, x_coords < p_ex)
            mask_y = tf.logical_and(y_coords >= p_sy, y_coords < p_ey)
            mask = tf.cast(tf.expand_dims(tf.logical_and(mask_x, mask_y), -1), dtype=img.dtype)
            
            # 将子图填充到对应位置,与原图合并
            padded_sub_img = tf.pad(
                sub_img,
                [[p_sx, height - p_ex], [p_sy, width - p_ey], [0, 0]]
            )
            return img * (1 - mask) + padded_sub_img * mask
            
        augmented_images = tf.map_fn(
            paste_single_image,
            (image, sub_images, paste_start_x, paste_start_y, paste_end_x, paste_end_y),
            dtype=image.dtype
        )
        
        return augmented_images

模型中调用示例

class Contrastive_learning_model(keras.Model):
    def __init__(self, cut_paste_augmentation):
        super().__init__()
        self.cut_paste = Cut_Paste(**cut_paste_augmentation)

    def train_step(self, data):
        images, labels = data  # 假设data包含图像和标签
        augmented_images_2 = self.cut_paste(images, training=True)
        
        # 后续训练逻辑...
        return super().train_step(data)

关键修改说明

  • 随机数生成:用tf.random.uniform替代random.randint,支持批量生成独立随机坐标,兼容TensorFlow图模式。
  • 张量修改方式:通过生成掩码+张量运算实现区域替换,避免直接赋值操作,符合TensorFlow的不可变张量特性。
  • 批量处理:使用tf.map_fn遍历批量中的每个样本,确保每个图像都有独立的裁剪/粘贴位置。
  • 训练/推理区分:添加training参数,仅在训练阶段应用增强,推理时返回原始图像。

内容的提问来源于stack exchange,提问作者Divyaj Shah

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 04:40:29