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

如何在Keras框架中实现GlobalMinPool2D自定义层?

Keras 自定义GlobalMinPool2D层实现方案

全局二维最小池化的核心逻辑是在输入特征图的空间维度(高、宽)上计算最小值,仅保留批量维度和通道维度,和官方提供的GlobalAvgPool2D、GlobalMaxPool2D输出结构完全对齐。

以下是符合Keras层规范的完整实现:

import tensorflow as tf
from tensorflow.keras import layers

class GlobalMinPool2D(layers.Layer):
    def __init__(self, data_format=None, keepdims=False, **kwargs):
        super().__init__(**kwargs)
        # 标准化数据格式,兼容channels_last和channels_first
        self.data_format = layers.normalize_data_format(data_format)
        self.keepdims = keepdims

    def call(self, inputs):
        if self.data_format == "channels_last":
            # 输入shape: (batch_size, height, width, channels),在高和宽维度取最小值
            return tf.reduce_min(inputs, axis=[1, 2], keepdims=self.keepdims)
        else:
            # 输入shape: (batch_size, channels, height, width),在高和宽维度取最小值
            return tf.reduce_min(inputs, axis=[2, 3], keepdims=self.keepdims)

    def get_config(self):
        # 实现序列化方法,支持模型保存加载
        config = super().get_config()
        config.update({
            "data_format": self.data_format,
            "keepdims": self.keepdims
        })
        return config

参数说明

  • data_format:和官方池化层参数逻辑完全一致,默认和Keras全局配置保持一致,可选值:
    • channels_last:输入形状为(batch_size, height, width, channels)
    • channels_first:输入形状为(batch_size, channels, height, width)
  • keepdims:布尔值,是否保留被降维的空间维度。设为True时,输出的空间维度会保留为1,方便后续和其他特征图做拼接等操作,默认值为False。

使用示例

# 构造测试输入:2个样本,32x32分辨率,16通道
inputs = tf.random.normal(shape=(2, 32, 32, 16))

# 基础使用
min_pool_layer = GlobalMinPool2D()
output = min_pool_layer(inputs)
print(output.shape) # 输出:(2, 16),和官方全局池化层输出结构一致

# 保留空间维度使用
min_pool_layer_keepdims = GlobalMinPool2D(keepdims=True)
output_keepdims = min_pool_layer_keepdims(inputs)
print(output_keepdims.shape) # 输出:(2, 1, 1, 16)

如果需要加载保存了该自定义层的模型,只需在load_model方法中传入custom_objects={"GlobalMinPool2D": GlobalMinPool2D}参数即可正常加载。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 01:06:04