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

TensorFlow实现图像cutout增强时张量赋值与tfa调用报错求解

问题原因及解决方法

1 自定义cutout函数报错解决

报错原因

  • TensorFlow的常规张量是不可变类型,没有assign方法,不能像numpy数组那样直接对切片赋值
  • 直接使用numpy的随机函数在tf.data的Graph执行模式下无法被追踪,会导致执行异常

修改后的兼容版本

def augment_cutout(image, label, size=68, n_squares=1):
    h = tf.shape(image)[0]
    w = tf.shape(image)[1]
    channels = tf.shape(image)[2]
    # 初始化掩码为全1
    mask = tf.ones((h, w, channels), dtype=image.dtype)
    for _ in range(n_squares):
        # 用TF的随机算子生成坐标
        y = tf.random.uniform(shape=[], minval=0, maxval=h, dtype=tf.int32)
        x = tf.random.uniform(shape=[], minval=0, maxval=w, dtype=tf.int32)
        size_half = size // 2
        y1 = tf.clip_by_value(y - size_half, 0, h)
        y2 = tf.clip_by_value(y + size_half, 0, h)
        x1 = tf.clip_by_value(x - size_half, 0, w)
        x2 = tf.clip_by_value(x + size_half, 0, w)
        # 生成cut区域的全0掩码
        cut_mask = tf.concat([
            tf.ones((y1, w, channels), dtype=image.dtype),
            tf.concat([
                tf.ones((y2-y1, x1, channels), dtype=image.dtype),
                tf.zeros((y2-y1, x2-x1, channels), dtype=image.dtype),
                tf.ones((y2-y1, w-x2, channels), dtype=image.dtype)
            ], axis=1),
            tf.ones((h-y2, w, channels), dtype=image.dtype)
        ], axis=0)
        mask = mask * cut_mask
    # 原图乘掩码得到cutout后的结果
    new_image = image * mask
    return tf.cast(new_image, tf.float32), label

调用位置要放在batch操作之前:

train_dataset = train_ds.map(map_func = preprocess_img, num_parallel_calls=AUTOTUNE)
# 先做cutout增强,再batch
train_dataset = train_dataset.map(map_func = augment_cutout, num_parallel_calls=AUTOTUNE)
train_dataset = train_dataset.shuffle(buffer_size=1000).batch(batch_size=BATCH_SIZE).prefetch(buffer_size=AUTOTUNE)

2 tfa.image.random_cutout报错解决

报错原因

tfa.image.random_cutout的输入要求是4维批量张量(batch_size, height, width, channels),你当前是在batch操作之前调用,输入是单张图片的3维张量(height, width, channels),接口读取维度时索引越界。

修复方案

把random_cut_out的调用放在batch操作之后即可:

import tensorflow_addons as tfa

def random_cut_out(images, labels):
    return tfa.image.random_cutout(images, (64, 64), constant_values = 1), labels

train_dataset = train_ds.map(map_func = preprocess_img, num_parallel_calls=AUTOTUNE) 
train_dataset = train_dataset.map(map_func = augment, num_parallel_calls=AUTOTUNE)
# 先batch,再调用cutout接口
train_dataset = train_dataset.shuffle(buffer_size=1000).batch(batch_size=BATCH_SIZE)
train_dataset = train_dataset.map(random_cut_out, num_parallel_calls=AUTOTUNE)
train_dataset = train_dataset.prefetch(buffer_size=AUTOTUNE)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 11:36:01