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

使用Tensorflow统计像素邻域内的像素颜色频次

问题描述

有形状为[1024, 1024, 3]的图像,颜色空间已缩减至约100种颜色。需要统计每一对颜色中,一种颜色出现在另一种颜色邻域内的次数(即邻域颜色直方图NCH),要求用TensorFlow实现并行计算,避免循环,支持通用N×N邻域大小及批量处理。

示例说明:
假设颜色空间含4种颜色(1、2、3、4),4×4图像如下:

[ [2 4 1 4]  
  [2 4 2 1]  
  [2 4 3 1]  
  [2 1 1 2] ]

对应的NCH矩阵(3×3邻域):

[ [8 7 4 6]  
  [7 6 2 12]  
  [4 2 0 2]  
  [6 12 2 4] ]

矩阵中NCH[i][j]表示颜色j出现在颜色i的邻域内的总次数。

已实现批量图像的N×N分块:

patches = tf.image.extract_patches(image, sizes=(1, D_size, D_size, 1), strides=(1, 1, 1, 1),
                                         padding='SAME', rates=[1, 1, 1, 1])
解决方案

步骤1:预处理输入图像

确保输入为单通道的颜色索引图(将RGB颜色映射为0~C-1的整数,C为颜色总数),批量输入形状为[batch_size, H, W, 1]。

步骤2:提取中心像素颜色

每个分块的中心像素对应原图像的像素,直接提取即可:

center_colors = tf.squeeze(image, axis=-1)  # 形状 [batch_size, H, W]

步骤3:分块展平与颜色计数

将每个N×N分块展平为一维向量,统计每个分块内各颜色的出现次数:

# 调整分块形状为 [batch_size, H, W, D_size*D_size]
patches_flat = tf.reshape(patches, [tf.shape(patches)[0], tf.shape(patches)[1], tf.shape(patches)[2], D_size*D_size])
# 对每个分块统计颜色出现次数,得到 [batch_size, H, W, C] 的直方图
color_counts = tf.math.bincount(patches_flat, minlength=C, maxlength=C, axis=-1)

步骤4:按中心颜色聚合计数

通过one-hot编码将中心颜色与对应分块的颜色计数关联,最终累加得到NCH矩阵:

# 将中心颜色转为one-hot编码,形状 [batch_size, H, W, C]
center_one_hot = tf.one_hot(center_colors, depth=C)
# 转置后做矩阵乘法,再在空间维度求和,得到批量的NCH矩阵 [batch_size, C, C]
nch = tf.matmul(tf.transpose(center_one_hot, perm=[0, 3, 1, 2]), color_counts)
nch = tf.reduce_sum(nch, axis=[2, 3])

完整可运行代码

import tensorflow as tf

def compute_nch(image, D_size, num_colors):
    # image: 批量输入,形状 [batch_size, H, W, 1],元素为0~num_colors-1的颜色索引
    # D_size: 邻域大小(N×N)
    # num_colors: 颜色总数
    
    # 提取N×N分块
    patches = tf.image.extract_patches(
        image,
        sizes=(1, D_size, D_size, 1),
        strides=(1, 1, 1, 1),
        padding='SAME',
        rates=[1, 1, 1, 1]
    )
    
    # 展平分块
    patches_flat = tf.reshape(patches, [tf.shape(patches)[0], tf.shape(patches)[1], tf.shape(patches)[2], D_size*D_size])
    
    # 每个分块的颜色计数
    color_counts = tf.math.bincount(patches_flat, minlength=num_colors, maxlength=num_colors, axis=-1)
    
    # 获取中心像素颜色
    center_colors = tf.squeeze(image, axis=-1)
    
    # 按中心颜色聚合,得到NCH矩阵
    center_one_hot = tf.one_hot(center_colors, depth=num_colors)
    nch = tf.matmul(tf.transpose(center_one_hot, perm=[0, 3, 1, 2]), color_counts)
    nch = tf.reduce_sum(nch, axis=[2, 3])
    
    return nch

# 测试示例
if __name__ == "__main__":
    # 将示例颜色1~4转为0~3的索引
    test_image = tf.constant([
        [2,4,1,4],
        [2,4,2,1],
        [2,4,3,1],
        [2,1,1,2]
    ], dtype=tf.int32) - 1
    # 调整为批量输入形状 [1,4,4,1]
    test_image = tf.expand_dims(tf.expand_dims(test_image, 0), -1)
    
    nch = compute_nch(test_image, D_size=3, num_colors=4)
    print("NCH矩阵:")
    print(nch.numpy()[0])

关键说明

  • 全程使用TensorFlow向量化操作,无显式循环,天然支持并行计算与批量处理。
  • 若输入为RGB图像,需先通过聚类(如K-Means)或颜色映射将其转为单通道的颜色索引图。
  • SAME padding确保每个像素都有对应的N×N邻域,若需边缘无填充可改用VALID padding。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 01:15:39