如何在TensorFlow自定义层中修改图像(附可运行示例)
效果对比
| 输入 | 预期输出 |
|---|---|
原始输入图像:![]() | 绘制黑色填充矩形后的预期输出图像:![]() |
报错原因
触发'tensorflow.python.framework.ops.EagerTensor' object has no attribute '__array_interface__'报错的核心原因是:PIL库的Image.fromarray方法仅支持接收实现了__array_interface__接口的numpy数组对象,无法直接解析传入的TensorFlow EagerTensor类型。
额外注意:不推荐在Keras自定义层的call方法中混用PIL、numpy等非TensorFlow算子,这类算子无法被TensorFlow编译为计算图,会导致模型无法导出为SavedModel格式、在tf.data流水线批量处理、使用XLA加速时出现兼容问题,运行性能也远低于原生TF算子。
推荐实现:纯TensorFlow原生算子实现(全场景兼容)
直接通过张量切片、拼接操作实现矩形区域填充,无额外格式转换开销,兼容eager模式、图模式、batch输入、模型序列化导出。
完整可运行代码如下:
import matplotlib.pyplot as plt import numpy as np import tensorflow as tf from PIL import Image class RemovePatch(tf.keras.layers.Layer): def __init__(self, x1=50, y1=50, x2=100, y2=100, fill_value=0, **kwargs): """ 绘制填充矩形的自定义数据增强层 :param x1: 矩形左上角x坐标(宽方向) :param y1: 矩形左上角y坐标(高方向) :param x2: 矩形右下角x坐标(宽方向) :param y2: 矩形右下角y坐标(高方向) :param fill_value: 矩形填充值,uint8格式下0为黑色、255为白色;浮点0-1格式下0为黑、1为白 """ super().__init__(**kwargs) self.x1 = x1 self.y1 = y1 self.x2 = x2 self.y2 = y2 self.fill_value = fill_value def call(self, image, training=None): if not training: return image img_shape = tf.shape(image) # 适配单张图(H,W,C)和batch输入(B,H,W,C) if len(image.shape) == 3: h, w, c = img_shape[0], img_shape[1], img_shape[2] # 坐标边界裁剪,防止矩形超出图像尺寸 y1 = tf.clip_by_value(self.y1, 0, h) y2 = tf.clip_by_value(self.y2, 0, h) x1 = tf.clip_by_value(self.x1, 0, w) x2 = tf.clip_by_value(self.x2, 0, w) # 拆分图像区域拼接填充块 top = image[:y1, :, :] patch_row_area = image[y1:y2, :, :] left = patch_row_area[:, :x1, :] fill_block = tf.fill((y2-y1, x2-x1, c), tf.cast(self.fill_value, image.dtype)) right = patch_row_area[:, x2:, :] mid = tf.concat([left, fill_block, right], axis=1) bottom = image[y2:, :, :] image = tf.concat([top, mid, bottom], axis=0) else: # batch维度输入逐张处理 outputs = tf.TensorArray(dtype=image.dtype, size=img_shape[0]) for i in tf.range(img_shape[0]): outputs = outputs.write(i, self.call(image[i], training=training)) image = outputs.stack() return image # 测试代码 layer = RemovePatch() image_file = "image.jpg" try: open(image_file) except FileNotFoundError: from requests import get r = get("https://picsum.photos/seed/picsum/300/300") with open(image_file, "wb") as f: f.write(r.content) with Image.open(image_file) as img: img = np.array(img) augmented = layer(img, training=True) augmented = np.array(augmented) plt.imshow(augmented) plt.show()
如果需要实现随机位置、随机大小的矩形擦除(Random Erasing常用增强逻辑),只需要在call方法中用tf.random.uniform生成随机的x1/y1/x2/y2坐标、随机填充值即可,核心填充逻辑不需要修改。
临时兼容PIL的实现方案(仅调试用,不推荐生产环境)
如果必须使用PIL的绘图能力(比如绘制复杂多边形、调用PIL内置滤镜),只需要先将EagerTensor转为numpy数组再传入PIL,处理完成后转回TensorFlow张量即可。注意该方案无法兼容图模式和模型导出,仅适合eager模式下临时调试。
核心修改代码如下:
def call(self, image, training=None): if not training: return image # EagerTensor转numpy数组 img_np = image.numpy() image_pil = Image.fromarray(img_np) ImageDraw.Draw(image_pil).rectangle([50, 50, 100, 100], fill="#000000") # 处理完成后转回TF张量 image = tf.convert_to_tensor(np.array(image_pil), dtype=image.dtype) return image
内容的提问来源于stack exchange,提问作者fejyesynb



