每日算法 — 使用java实现K-means聚类:图像颜色量化与肘部法则选K

K-means聚类是机器学习领域最经典的无监督学习算法之一,其核心思想是通过迭代优化将数据点划分为K个簇,使得簇内平方和最小。本文以图像颜色量化为实战场景,用Java从零实现完整的K-means算法,讲解随机初始化、均值更新、簇分配与肘部法则确定最佳K值,所有代码均可直接运行。

一、算法原理与核心步骤

1.1 什么是K-means聚类

K-means的目标是将N个样本点划分为K个簇,使得每个样本点到其所属簇中心的距离之和最小。形式化地说,我们要最小化以下目标函数:

$$J = \sum_{i=1}^{K} \sum_{x \in C_i} ||x – \mu_i||^2$$

其中 $C_i$ 表示第 $i$ 个簇,$\mu_i$ 表示第 $i$ 个簇的中心点。由于该优化问题是NP-hard的,K-means采用贪心迭代策略寻找局部最优解。

1.2 标准K-means流程

算法遵循经典的”分配-更新”交替迭代框架:

  1. 初始化:随机选择K个样本点作为初始簇中心
  2. 分配步(Assignment):将每个样本点分配到距离最近的簇中心
  3. 更新步(Update):重新计算每个簇的均值,作为新的簇中心
  4. 收敛判断:若簇中心变化小于阈值或达到最大迭代次数,则停止

该过程保证每次迭代都会降低目标函数值,最终收敛到局部最优。

二、图像颜色量化的场景建模

2.1 为什么用K-means做颜色量化

一张RGB彩色图像的每个像素由三个8位通道组成,总共可表示约1677万种颜色。但在很多实际场景(如复古风格渲染、图标生成、低带宽传输)中,我们只需要少数代表性颜色即可呈现图像主体。颜色量化正是将全彩图像压缩为K种颜色的过程。

将每个像素视为三维空间 $(R, G, B)$ 中的一个点,K-means天然适合将颜色空间聚为K个簇,每个簇中心即为代表色。

2.2 数据结构设计

为了代码的可读性与扩展性,我们定义三个核心类:ColorPoint(颜色点)、Cluster(簇)、KMeans(算法引擎)。

import java.awt.image.BufferedImage;
import java.io.File;
import java.io.IOException;
import java.util.*;
import javax.imageio.ImageIO;

/**
 * 颜色点:将像素的RGB值映射为三维空间中的点
 * 每个分量归一化到[0,1]区间,便于距离计算与可视化
 */
class ColorPoint {
    final double r, g, b;
    final int x, y; // 像素在图像中的原始坐标(可选,用于调试)

    ColorPoint(double r, double g, double b) {
        this(r, g, b, -1, -1);
    }

    ColorPoint(double r, double g, double b, int x, int y) {
        this.r = r;
        this.g = g;
        this.b = b;
        this.x = x;
        this.y = y;
    }

    /**
     * 计算与另一个颜色点的欧几里得距离
     * 在RGB空间中,欧氏距离近似人眼感知的颜色差异
     */
    double distanceTo(ColorPoint other) {
        double dr = this.r - other.r;
        double dg = this.g - other.g;
        double db = this.b - other.b;
        return Math.sqrt(dr * dr + dg * dg + db * db);
    }

    /**
     * 计算与另一个颜色点的平方距离
     * 避免开方运算,仅在比较距离大小时使用,提升性能
     */
    double squaredDistanceTo(ColorPoint other) {
        double dr = this.r - other.r;
        double dg = this.g - other.g;
        double db = this.b - other.b;
        return dr * dr + dg * dg + db * db;
    }

    /**
     * 将归一化颜色值转换为0-255的整数,用于生成量化后的图像
     */
    int toRgbInt() {
        int ri = (int) Math.round(r * 255.0);
        int gi = (int) Math.round(g * 255.0);
        int bi = (int) Math.round(b * 255.0);
        ri = Math.min(255, Math.max(0, ri));
        gi = Math.min(255, Math.max(0, gi));
        bi = Math.min(255, Math.max(0, bi));
        return (ri << 16) | (gi << 8) | bi;
    }

    @Override
    public String toString() {
        return String.format("RGB(%.3f, %.3f, %.3f)", r, g, b);
    }
}

/**
 * 簇:维护一个簇的中心点以及属于该簇的所有颜色点
 * 在每次迭代中,簇中心会根据成员点的均值重新计算
 */
class Cluster {
    ColorPoint centroid;   // 当前簇中心
    List<ColorPoint> points; // 当前分配到此簇的点

    Cluster(ColorPoint centroid) {
        this.centroid = centroid;
        this.points = new ArrayList<>();
    }

    /**
     * 清空当前簇的成员列表,为下一轮分配做准备
     */
    void clearPoints() {
        points.clear();
    }

    /**
     * 将一个颜色点加入本簇
     */
    void addPoint(ColorPoint p) {
        points.add(p);
    }

    /**
     * 根据当前成员点重新计算簇中心(均值)
     * 若簇为空,则保持原中心不变(防止除零)
     * @return 新的簇中心
     */
    ColorPoint computeNewCentroid() {
        if (points.isEmpty()) {
            return centroid;
        }
        double sumR = 0, sumG = 0, sumB = 0;
        for (ColorPoint p : points) {
            sumR += p.r;
            sumG += p.g;
            sumB += p.b;
        }
        return new ColorPoint(sumR / points.size(), sumG / points.size(), sumB / points.size());
    }

    /**
     * 计算簇内平方和(Within-Cluster Sum of Squares, WCSS)
     * 用于评估聚类紧密度和肘部法则
     */
    double computeWcss() {
        double sum = 0;
        for (ColorPoint p : points) {
            sum += p.squaredDistanceTo(centroid);
        }
        return sum;
    }
}

三、K-means算法引擎

3.1 K-Means++初始化策略

标准K-means采用完全随机初始化,容易选到相近的初始中心,导致收敛到较差的局部最优。K-Means++通过概率抽样改善初始化质量:第一个中心随机选,后续中心以与已有中心距离的平方成正比的概率选取,使得初始中心尽可能分散。

/**
 * K-means聚类引擎
 * 支持标准随机初始化和K-Means++优化初始化
 */
class KMeans {
    private final int k;                  // 簇的数量
    private final int maxIterations;      // 最大迭代次数
    private final double tolerance;       // 中心点变化收敛阈值
    private List<Cluster> clusters;       // 簇列表
    private final Random random;

    KMeans(int k, int maxIterations, double tolerance, long seed) {
        this.k = k;
        this.maxIterations = maxIterations;
        this.tolerance = tolerance;
        this.random = new Random(seed);
    }

    /**
     * 主入口:对输入的颜色点集合执行K-means聚类
     * @param points 所有像素对应的颜色点
     * @return 聚类后的簇列表
     */
    List<Cluster> fit(List<ColorPoint> points) {
        if (points.size() < k) {
            throw new IllegalArgumentException("点数必须大于等于K值");
        }

        // 使用K-Means++初始化簇中心
        initializePlusPlus(points);

        for (int iter = 0; iter < maxIterations; iter++) {
            // 步骤1:清空各簇成员
            for (Cluster c : clusters) {
                c.clearPoints();
            }

            // 步骤2:分配每个点到最近的簇
            for (ColorPoint p : points) {
                Cluster nearest = findNearestCluster(p);
                nearest.addPoint(p);
            }

            // 步骤3:更新簇中心并检查收敛
            double maxShift = 0;
            for (Cluster c : clusters) {
                ColorPoint oldCentroid = c.centroid;
                ColorPoint newCentroid = c.computeNewCentroid();
                double shift = oldCentroid.distanceTo(newCentroid);
                maxShift = Math.max(maxShift, shift);
                c.centroid = newCentroid;
            }

            if (maxShift < tolerance) {
                System.out.println("收敛于第 " + (iter + 1) + " 轮,最大中心偏移: " + String.format("%.6f", maxShift));
                break;
            }
        }

        return clusters;
    }

    /**
     * K-Means++初始化
     * 1. 随机选一个初始中心
     * 2. 对每个点,计算它到最近已选中心的距离D(x)
     * 3. 以概率 D(x)^2 / sum(D^2) 选择下一个中心
     * 4. 重复直到选够K个中心
     */
    private void initializePlusPlus(List<ColorPoint> points) {
        clusters = new ArrayList<>();
        List<ColorPoint> centroids = new ArrayList<>();

        // 第1个中心:完全随机
        centroids.add(points.get(random.nextInt(points.size())));

        // 第2~K个中心:按距离平方加权抽样
        while (centroids.size() < k) {
            double[] minDistSq = new double[points.size()];
            double totalDistSq = 0;

            for (int i = 0; i < points.size(); i++) {
                double minSq = Double.MAX_VALUE;
                for (ColorPoint c : centroids) {
                    double d = points.get(i).squaredDistanceTo(c);
                    if (d < minSq) minSq = d;
                }
                minDistSq[i] = minSq;
                totalDistSq += minSq;
            }

            // 轮盘赌选择
            double threshold = random.nextDouble() * totalDistSq;
            double cumulative = 0;
            int selected = 0;
            for (int i = 0; i < points.size(); i++) {
                cumulative += minDistSq[i];
                if (cumulative >= threshold) {
                    selected = i;
                    break;
                }
            }
            centroids.add(points.get(selected));
        }

        for (ColorPoint c : centroids) {
            clusters.add(new Cluster(c));
        }
    }

    /**
     * 查找距离指定点最近的簇
     */
    private Cluster findNearestCluster(ColorPoint p) {
        Cluster nearest = clusters.get(0);
        double minDist = p.squaredDistanceTo(nearest.centroid);

        for (int i = 1; i < clusters.size(); i++) {
            double d = p.squaredDistanceTo(clusters.get(i).centroid);
            if (d < minDist) {
                minDist = d;
                nearest = clusters.get(i);
            }
        }
        return nearest;
    }

    /**
     * 计算所有簇的WCSS之和
     */
    double computeTotalWcss() {
        double total = 0;
        for (Cluster c : clusters) {
            total += c.computeWcss();
        }
        return total;
    }

    /**
     * 获取聚类完成后的簇中心列表
     */
    List<ColorPoint> getCentroids() {
        List<ColorPoint> list = new ArrayList<>();
        for (Cluster c : clusters) {
            list.add(c.centroid);
        }
        return list;
    }
}

四、肘部法则:如何选出最佳K值

4.1 肘部法则原理

K-means需要预先指定K值,但K值的选择直接影响聚类效果。一种直观的方法是尝试多个K值,计算对应的总WCSS(Within-Cluster Sum of Squares),绘制K-WCSS曲线。当K增加到某个值后,WCSS的下降幅度会显著减缓,曲线在此处形成”手肘”形状,对应的K即为较优选择。

/**
 * 肘部法则分析器
 * 对K从1到maxK分别运行多次K-means,取平均WCSS绘制肘部曲线
 */
class ElbowAnalyzer {
    private final int maxK;
    private final int repeats;        // 每个K重复实验次数(降低随机性影响)
    private final int maxIterations;
    private final double tolerance;

    ElbowAnalyzer(int maxK, int repeats, int maxIterations, double tolerance) {
        this.maxK = maxK;
        this.repeats = repeats;
        this.maxIterations = maxIterations;
        this.tolerance = tolerance;
    }

    /**
     * 执行肘部分析
     * @param points 颜色点集合
     * @return 每个K对应的平均WCSS数组,索引0对应K=1
     */
    double[] analyze(List<ColorPoint> points) {
        double[] avgWcss = new double[maxK];
        System.out.println("\n========== 肘部法则分析 ==========");
        System.out.println("K\t平均WCSS\t\t下降率");
        System.out.println("----------------------------------");

        for (int k = 1; k <= maxK; k++) {
            double sumWcss = 0;
            for (int r = 0; r < repeats; r++) {
                KMeans km = new KMeans(k, maxIterations, tolerance, System.nanoTime() + r);
                km.fit(points);
                sumWcss += km.computeTotalWcss();
            }
            avgWcss[k - 1] = sumWcss / repeats;

            double dropRate = 0;
            if (k > 1) {
                dropRate = (avgWcss[k - 2] - avgWcss[k - 1]) / avgWcss[k - 2] * 100;
            }
            System.out.printf("%d\t%.4f\t\t%.2f%%%n", k, avgWcss[k - 1], dropRate);
        }
        return avgWcss;
    }

    /**
     * 自动推荐最佳K值(简化版:找下降率首次低于阈值的位置)
     */
    int suggestK(double[] avgWcss, double thresholdRate) {
        for (int k = 2; k < avgWcss.length; k++) {
            double dropRate = (avgWcss[k - 2] - avgWcss[k - 1]) / avgWcss[k - 2];
            if (dropRate < thresholdRate) {
                return k;
            }
        }
        return avgWcss.length;
    }
}

五、图像读取与量化输出

5.1 工具类:图像与颜色的互转

/**
 * 图像处理工具类
 * 负责加载图像、提取颜色点、生成量化后的输出图像
 */
class ImageUtils {

    /**
     * 从文件加载图像,并将每个像素提取为ColorPoint
     * 支持降采样以加速聚类(每step个像素取一个)
     */
    static List<ColorPoint> loadImagePoints(String path, int step) throws IOException {
        BufferedImage image = ImageIO.read(new File(path));
        if (image == null) {
            throw new IOException("无法读取图像: " + path);
        }
        List<ColorPoint> points = new ArrayList<>();
        int width = image.getWidth();
        int height = image.getHeight();

        for (int y = 0; y < height; y += step) {
            for (int x = 0; x < width; x += step) {
                int rgb = image.getRGB(x, y);
                double r = ((rgb >> 16) & 0xFF) / 255.0;
                double g = ((rgb >> 8) & 0xFF) / 255.0;
                double b = (rgb & 0xFF) / 255.0;
                points.add(new ColorPoint(r, g, b, x, y));
            }
        }
        System.out.println("图像尺寸: " + width + "x" + height + ", 采样点数: " + points.size());
        return points;
    }

    /**
     * 根据聚类结果生成颜色量化后的图像
     * 每个像素用其所属簇的中心颜色替换
     */
    static BufferedImage quantizeImage(BufferedImage original, List<Cluster> clusters, int step) {
        int width = original.getWidth();
        int height = original.getHeight();
        BufferedImage output = new BufferedImage(width, height, BufferedImage.TYPE_INT_RGB);

        // 为每个采样像素建立映射,非采样像素用最近采样像素的簇结果
        int[][] clusterMap = new int[height][width];
        for (Cluster cluster : clusters) {
            int centroidRgb = cluster.centroid.toRgbInt();
            for (ColorPoint p : cluster.points) {
                if (p.x >= 0 && p.y >= 0) {
                    clusterMap[p.y][p.x] = centroidRgb;
                }
            }
        }

        // 填充整幅图像:对未采样的像素,找最近的采样点所属的簇颜色
        for (int y = 0; y < height; y++) {
            for (int x = 0; x < width; x++) {
                int nearestRgb = findNearestSampledColor(x, y, clusterMap, step);
                output.setRGB(x, y, nearestRgb);
            }
        }
        return output;
    }

    /**
     * 为非采样像素查找最近采样像素的颜色值
     */
    private static int findNearestSampledColor(int x, int y, int[][] clusterMap, int step) {
        int sy = (y / step) * step;
        int sx = (x / step) * step;
        // 在采样网格中找最近点
        int bestY = Math.min(sy, clusterMap.length - 1);
        int bestX = Math.min(sx, clusterMap[0].length - 1);
        return clusterMap[bestY][bestX];
    }

    /**
     * 保存图像到文件
     */
    static void saveImage(BufferedImage image, String path) throws IOException {
        File out = new File(path);
        ImageIO.write(image, "png", out);
        System.out.println("已保存量化图像: " + out.getAbsolutePath());
    }

    /**
     * 将原始图像完整转换为ColorPoint列表(无降采样,用于最终精确量化)
     */
    static List<ColorPoint> loadAllPixels(BufferedImage image) {
        List<ColorPoint> points = new ArrayList<>();
        int width = image.getWidth();
        int height = image.getHeight();
        for (int y = 0; y < height; y++) {
            for (int x = 0; x < width; x++) {
                int rgb = image.getRGB(x, y);
                double r = ((rgb >> 16) & 0xFF) / 255.0;
                double g = ((rgb >> 8) & 0xFF) / 255.0;
                double b = (rgb & 0xFF) / 255.0;
                points.add(new ColorPoint(r, g, b, x, y));
            }
        }
        return points;
    }
}

六、主程序与运行演示

6.1 完整入口

/**
 * K-means图像颜色量化演示程序
 * 用法:提供一张输入图像路径,程序先执行肘部分析推荐K值,
 * 然后用该K值对图像进行颜色量化并输出结果。
 */
public class KMeansColorQuantization {
    public static void main(String[] args) throws IOException {
        String inputPath = args.length > 0 ? args[0] : "input.png";
        String outputPath = args.length > 1 ? args[1] : "output_quantized.png";

        // 参数配置
        final int MAX_K_FOR_ELBOW = 10;      // 肘部分析的最大K值
        final int ELBOW_REPEATS = 3;         // 每个K重复实验次数
        final int MAX_ITERATIONS = 100;      // K-means最大迭代轮数
        final double TOLERANCE = 1e-4;       // 收敛阈值
        final int SAMPLE_STEP = 4;           // 降采样步长(4表示每4x4像素取1个样本)

        // 1. 加载图像并降采样(加速肘部分析)
        System.out.println("=== 加载图像 ===");
        List<ColorPoint> sampledPoints = ImageUtils.loadImagePoints(inputPath, SAMPLE_STEP);

        if (sampledPoints.isEmpty()) {
            System.err.println("图像为空或读取失败");
            return;
        }

        // 2. 肘部法则分析推荐最佳K值
        ElbowAnalyzer analyzer = new ElbowAnalyzer(MAX_K_FOR_ELBOW, ELBOW_REPEATS, MAX_ITERATIONS, TOLERANCE);
        double[] wcssCurve = analyzer.analyze(sampledPoints);
        int suggestedK = analyzer.suggestK(wcssCurve, 0.15); // 下降率低于15%视为肘部
        System.out.println("推荐K值: " + suggestedK);

        // 3. 使用推荐K值重新加载完整像素并执行最终聚类
        System.out.println("\n=== 执行最终颜色量化 (K=" + suggestedK + ") ===");
        BufferedImage original = ImageIO.read(new File(inputPath));
        List<ColorPoint> allPoints = ImageUtils.loadAllPixels(original);

        KMeans finalKMeans = new KMeans(suggestedK, MAX_ITERATIONS, TOLERANCE, 42);
        List<Cluster> finalClusters = finalKMeans.fit(allPoints);

        double finalWcss = finalKMeans.computeTotalWcss();
        System.out.println("最终WCSS: " + String.format("%.4f", finalWcss));
        System.out.println("簇中心颜色:");
        for (int i = 0; i < finalClusters.size(); i++) {
            ColorPoint c = finalClusters.get(i).centroid;
            System.out.printf("  簇%d: RGB(%d, %d, %d)  像素数: %d%n",
                i + 1,
                (int) Math.round(c.r * 255),
                (int) Math.round(c.g * 255),
                (int) Math.round(c.b * 255),
                finalClusters.get(i).points.size());
        }

        // 4. 生成并保存量化图像
        BufferedImage quantized = ImageUtils.quantizeImage(original, finalClusters, 1);
        ImageUtils.saveImage(quantized, outputPath);
        System.out.println("\n颜色量化完成!");
    }
}

6.2 运行效果说明

由于当前环境为纯控制台,读者可将上述代码保存为 KMeansColorQuantization.java,准备一张PNG图像后执行:

javac KMeansColorQuantization.java
java KMeansColorQuantization input.png output.png

程序将依次输出:
1. 图像尺寸与采样点数
2. K从1到10的WCSS曲线与下降率
3. 推荐K值与最终聚类的簇中心颜色
4. 量化后的PNG文件

以一张自然风景照为例,当K=8时,通常WCSS下降率已从K=2时的60%降至10%以下,说明8种颜色已能较好地概括图像色彩分布。

七、算法复杂度分析

步骤 时间复杂度 空间复杂度 关键瓶颈
K-Means++初始化 O(K·N) O(K) 需计算所有点到最近中心的距离
分配步(每轮) O(K·N) O(N) 每个点需与K个中心比较
更新步(每轮) O(N) O(K) 遍历各簇成员求均值
单次完整运行 O(I·K·N) O(N+K) I为迭代次数,通常I<50
肘部法则 O(R·K_max·I·K·N) O(N+K) R为重复次数,可用降采样大幅降低N

其中N为像素数。对于1080P图像(约200万像素),直接聚类较慢;通过降采样(step=4将N降至12万)可在秒级完成肘部分析,最终用完整像素跑一次精修即可。

八、扩展方向

  • Lab颜色空间:RGB空间中欧氏距离与人眼感知不完全一致,可将点转换到CIE Lab空间再聚类,颜色量化结果更自然。
  • 超像素分割(SLIC):将K-means从颜色维度扩展到空间维度,同时考虑像素的位置与颜色,生成边界更规整的超像素。
  • K-medoids与核K-means:当存在噪声点时,K-medoids用实际样本点作为中心更鲁棒;核方法可将数据映射到高维空间处理非球形簇。
  • 动画颜色迁移:将源图像的K个簇中心映射到目标图像的K个簇中心,实现整体色调的风格迁移。

K-means虽然简单,却是理解无监督学习的最佳起点。通过图像颜色量化这一直观场景,我们不仅实现了算法的每个核心环节,还引入了K-Means++和肘部法则这两个工程实践中不可或缺的优化手段。

发表回复

您的邮箱地址不会被公开。 必填项已用 * 标注