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

如何在TensorFlow/Keras模型中添加按比例缩放的调整大小层

TensorFlow/Keras按比例缩放图像层实现方案

你可以通过自定义Keras层的方式实现按比例缩放的图像调整层,完全支持要求的所有插值方法,同时兼容模型的训练、推理与序列化保存。

方案1:自定义可复用比例缩放层(推荐)

该实现符合Keras层规范,没有序列化兼容问题,生产环境优先使用:

import tensorflow as tf
from tensorflow.keras import layers

class ResizeByScale(layers.Layer):
    def __init__(self, scale_factor: float, interpolation: str = "bilinear", **kwargs):
        super().__init__(**kwargs)
        # 缩放比例,大于1为上采样,小于1为下采样
        self.scale_factor = scale_factor
        # 插值方法,和tf.image.resize支持的参数完全一致
        self.interpolation = interpolation

    def call(self, inputs):
        # 动态获取输入图像的高、宽,默认输入格式为NHWC(批量数、高、宽、通道数)
        input_shape = tf.shape(inputs)
        new_height = tf.cast(tf.cast(input_shape[1], tf.float32) * self.scale_factor, tf.int32)
        new_width = tf.cast(tf.cast(input_shape[2], tf.float32) * self.scale_factor, tf.int32)
        return tf.image.resize(inputs, size=(new_height, new_width), method=self.interpolation)

    def get_config(self):
        # 实现配置方法,支持模型保存与加载
        config = super().get_config()
        config.update({
            "scale_factor": self.scale_factor,
            "interpolation": self.interpolation
        })
        return config

使用示例

# 构建测试模型,输入shape为(None, 100, 100, 3)
inputs = layers.Input(shape=(100, 100, 3))
# 添加2倍上采样层,使用bicubic插值
x = ResizeByScale(scale_factor=2, interpolation="bicubic")(inputs)
model = tf.keras.Model(inputs=inputs, outputs=x)

# 验证输出尺寸
test_input = tf.random.normal((1, 100, 100, 3))
test_output = model(test_input)
print(test_output.shape) # 输出为 (1, 200, 200, 3),符合预期

方案2:Lambda层快速实现(适合临时测试)

如果仅做快速验证,也可以用Lambda层直接封装逻辑,缺点是模型保存可能存在兼容问题:

scale_factor = 2
interpolation = "nearest"

resize_by_scale_layer = layers.Lambda(
    lambda img: tf.image.resize(
        img,
        size=(
            tf.cast(tf.shape(img)[1] * scale_factor, tf.int32),
            tf.cast(tf.shape(img)[2] * scale_factor, tf.int32)
        ),
        method=interpolation
    )
)

支持的插值方法列表

两种实现都完全兼容tf.image.resize的所有插值参数,包含常用的四类插值:

  • bilinear:双线性插值
  • nearest:最近邻插值
  • bicubic:双三次插值
  • area:区域插值
  • 额外支持:gaussian、lanczos3、lanczos5、mitchellcubic

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 20:39:04