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流程
算法遵循经典的”分配-更新”交替迭代框架:
- 初始化:随机选择K个样本点作为初始簇中心
- 分配步(Assignment):将每个样本点分配到距离最近的簇中心
- 更新步(Update):重新计算每个簇的均值,作为新的簇中心
- 收敛判断:若簇中心变化小于阈值或达到最大迭代次数,则停止
该过程保证每次迭代都会降低目标函数值,最终收敛到局部最优。
二、图像颜色量化的场景建模
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++和肘部法则这两个工程实践中不可或缺的优化手段。