基于GMM与EM算法的RGB图像分割Java实现求助
我完全懂你找不到纯手写EM算法实现RGB图像分割资料的 frustration——大部分示例都是直接调用封装好的库,根本看不到底层逻辑。下面我给你一个完整的Java实现,基于高斯混合模型(GMM)的EM算法,每一步都有注释,帮你搞懂底层原理:
Java实现RGB图像分割的EM算法示例
核心思路
EM算法在这里分为两步循环执行,直到参数收敛:
- E步:计算每个像素属于每个高斯簇的后验概率(也就是“期望”每个簇对像素的贡献)
- M步:根据E步得到的概率,更新每个高斯簇的均值、协方差和权重(最大化似然估计)
对于RGB图像,每个像素是一个三维向量(R, G, B),我们假设这些向量服从K个高斯分布的混合模型。
完整Java代码示例
import java.awt.image.BufferedImage; import java.io.File; import java.io.IOException; import javax.imageio.ImageIO; import java.util.Random; public class RGBAEMSegmentation { // 高斯混合模型的簇数量,可根据需要调整 private static final int K = 3; // 迭代终止阈值,当参数变化小于这个值时停止 private static final double THRESHOLD = 1e-6; // 最大迭代次数 private static final int MAX_ITERATIONS = 100; // 存储每个高斯簇的参数:均值(R,G,B)、协方差矩阵、权重 private static double[][] means = new double[K][3]; private static double[][][] covariances = new double[K][3][3]; private static double[] weights = new double[K]; public static void main(String[] args) throws IOException { // 加载RGB图像 BufferedImage image = ImageIO.read(new File("input.jpg")); int width = image.getWidth(); int height = image.getHeight(); int totalPixels = width * height; // 提取所有像素的RGB数据,存储为三维数组[pixelIndex][R/G/B] double[][] pixels = new double[totalPixels][3]; int idx = 0; for (int y = 0; y < height; y++) { for (int x = 0; x < width; x++) { int rgb = image.getRGB(x, y); pixels[idx][0] = (rgb >> 16) & 0xFF; // R通道 pixels[idx][1] = (rgb >> 8) & 0xFF; // G通道 pixels[idx][2] = rgb & 0xFF; // B通道 idx++; } } // 初始化GMM参数 initializeGMM(pixels); // 迭代EM算法 double logLikelihood = Double.NEGATIVE_INFINITY; for (int iter = 0; iter < MAX_ITERATIONS; iter++) { // E步:计算每个像素属于每个簇的后验概率 double[][] posteriors = eStep(pixels); // M步:更新GMM参数 double newLogLikelihood = mStep(pixels, posteriors); // 检查是否收敛 if (Math.abs(newLogLikelihood - logLikelihood) < THRESHOLD) { System.out.println("迭代收敛,共执行" + (iter+1) + "次"); break; } logLikelihood = newLogLikelihood; } // 生成分割后的图像 BufferedImage segmentedImage = new BufferedImage(width, height, BufferedImage.TYPE_INT_RGB); idx = 0; for (int y = 0; y < height; y++) { for (int x = 0; x < width; x++) { // 找到当前像素所属的簇(概率最大的那个) int cluster = getPixelCluster(pixels[idx]); // 用簇的均值作为该区域的颜色(也可以用原始颜色,这里为了可视化更清晰) int r = (int) Math.round(means[cluster][0]); int g = (int) Math.round(means[cluster][1]); int b = (int) Math.round(means[cluster][2]); segmentedImage.setRGB(x, y, (r << 16) | (g << 8) | b); idx++; } } // 保存分割后的图像 ImageIO.write(segmentedImage, "jpg", new File("segmented_output.jpg")); System.out.println("分割完成,结果已保存为segmented_output.jpg"); } // 初始化GMM参数:随机选择K个像素作为初始均值,协方差初始为单位矩阵,权重初始为1/K private static void initializeGMM(double[][] pixels) { Random random = new Random(); int totalPixels = pixels.length; // 初始化权重 for (int k = 0; k < K; k++) { weights[k] = 1.0 / K; } // 初始化均值:随机选K个不同的像素 boolean[] selected = new boolean[totalPixels]; for (int k = 0; k < K; k++) { int idx; do { idx = random.nextInt(totalPixels); } while (selected[idx]); selected[idx] = true; means[k][0] = pixels[idx][0]; means[k][1] = pixels[idx][1]; means[k][2] = pixels[idx][2]; } // 初始化协方差矩阵:单位矩阵乘以一个常数(这里用255/4,避免初始方差太小) double initCov = (255.0 / 4) * (255.0 / 4); for (int k = 0; k < K; k++) { for (int i = 0; i < 3; i++) { for (int j = 0; j < 3; j++) { covariances[k][i][j] = (i == j) ? initCov : 0; } } } } // E步:计算每个像素属于每个簇的后验概率 private static double[][] eStep(double[][] pixels) { int totalPixels = pixels.length; double[][] posteriors = new double[totalPixels][K]; for (int i = 0; i < totalPixels; i++) { double[] pixel = pixels[i]; double sum = 0.0; // 先计算每个簇的似然乘以权重 for (int k = 0; k < K; k++) { double likelihood = calculateGaussian(pixel, means[k], covariances[k]); posteriors[i][k] = weights[k] * likelihood; sum += posteriors[i][k]; } // 归一化得到后验概率 for (int k = 0; k < K; k++) { posteriors[i][k] /= sum; } } return posteriors; } // 计算三维高斯分布的概率密度 private static double calculateGaussian(double[] x, double[] mu, double[][] sigma) { // 计算(x - mu) double[] diff = new double[3]; for (int i = 0; i < 3; i++) { diff[i] = x[i] - mu[i]; } // 计算协方差矩阵的逆和行列式 double[][] invSigma = matrixInverse(sigma); double detSigma = matrixDeterminant(sigma); // 高斯分布公式:(1/(sqrt((2π)^d * det(Sigma)))) * exp(-0.5*(x-mu)^T * inv(Sigma) * (x-mu)) double exponent = -0.5 * matrixDotProduct(matrixTranspose(new double[][]{diff}), invSigma)[0][0]; double denominator = Math.pow(2 * Math.PI, 3.0/2) * Math.sqrt(detSigma); return Math.exp(exponent) / denominator; } // M步:更新均值、协方差和权重 private static double mStep(double[][] pixels, double[][] posteriors) { int totalPixels = pixels.length; double logLikelihood = 0.0; for (int k = 0; k < K; k++) { // 计算当前簇的总权重(N_k) double nk = 0.0; for (int i = 0; i < totalPixels; i++) { nk += posteriors[i][k]; } // 更新权重 weights[k] = nk / totalPixels; // 更新均值 double[] newMean = new double[3]; for (int i = 0; i < totalPixels; i++) { for (int d = 0; d < 3; d++) { newMean[d] += posteriors[i][k] * pixels[i][d]; } } for (int d = 0; d < 3; d++) { newMean[d] /= nk; } means[k] = newMean; // 更新协方差矩阵 double[][] newCov = new double[3][3]; for (int i = 0; i < totalPixels; i++) { double[] diff = new double[3]; for (int d = 0; d < 3; d++) { diff[d] = pixels[i][d] - means[k][d]; } double[][] outerProduct = outerProduct(diff, diff); for (int row = 0; row < 3; row++) { for (int col = 0; col < 3; col++) { newCov[row][col] += posteriors[i][k] * outerProduct[row][col]; } } } for (int row = 0; row < 3; row++) { for (int col = 0; col < 3; col++) { newCov[row][col] /= nk; } } covariances[k] = newCov; } // 计算当前的对数似然 for (int i = 0; i < totalPixels; i++) { double sum = 0.0; for (int k = 0; k < K; k++) { sum += weights[k] * calculateGaussian(pixels[i], means[k], covariances[k]); } logLikelihood += Math.log(sum); } return logLikelihood; } // 获取像素所属的簇(概率最大的那个) private static int getPixelCluster(double[] pixel) { int bestCluster = 0; double maxProb = 0.0; for (int k = 0; k < K; k++) { double prob = weights[k] * calculateGaussian(pixel, means[k], covariances[k]); if (prob > maxProb) { maxProb = prob; bestCluster = k; } } return bestCluster; } // 辅助方法:计算矩阵的逆(仅针对3x3矩阵) private static double[][] matrixInverse(double[][] mat) { double det = matrixDeterminant(mat); if (Math.abs(det) < 1e-10) { // 如果行列式接近0,返回单位矩阵避免除以0 return new double[][]{{1,0,0}, {0,1,0}, {0,0,1}}; } double[][] adjugate = new double[3][3]; adjugate[0][0] = mat[1][1]*mat[2][2] - mat[1][2]*mat[2][1]; adjugate[0][1] = mat[0][2]*mat[2][1] - mat[0][1]*mat[2][2]; adjugate[0][2] = mat[0][1]*mat[1][2] - mat[0][2]*mat[1][1]; adjugate[1][0] = mat[1][2]*mat[2][0] - mat[1][0]*mat[2][2]; adjugate[1][1] = mat[0][0]*mat[2][2] - mat[0][2]*mat[2][0]; adjugate[1][2] = mat[0][2]*mat[1][0] - mat[0][0]*mat[1][2]; adjugate[2][0] = mat[1][0]*mat[2][1] - mat[1][1]*mat[2][0]; adjugate[2][1] = mat[0][1]*mat[2][0] - mat[0][0]*mat[2][1]; adjugate[2][2] = mat[0][0]*mat[1][1] - mat[0][1]*mat[1][0]; double[][] inv = new double[3][3]; for (int i = 0; i < 3; i++) { for (int j = 0; j < 3; j++) { inv[i][j] = adjugate[i][j] / det; } } return inv; } // 辅助方法:计算3x3矩阵的行列式 private static double matrixDeterminant(double[][] mat) { return mat[0][0]*(mat[1][1]*mat[2][2] - mat[1][2]*mat[2][1]) - mat[0][1]*(mat[1][0]*mat[2][2] - mat[1][2]*mat[2][0]) + mat[0][2]*(mat[1][0]*mat[2][1] - mat[1][1]*mat[2][0]); } // 辅助方法:计算矩阵转置 private static double[][] matrixTranspose(double[][] mat) { int rows = mat.length; int cols = mat[0].length; double[][] transposed = new double[cols][rows]; for (int i = 0; i < rows; i++) { for (int j = 0; j < cols; j++) { transposed[j][i] = mat[i][j]; } } return transposed; } // 辅助方法:计算两个矩阵的点积 private static double[][] matrixDotProduct(double[][] a, double[][] b) { int aRows = a.length; int aCols = a[0].length; int bCols = b[0].length; double[][] result = new double[aRows][bCols]; for (int i = 0; i < aRows; i++) { for (int j = 0; j < bCols; j++) { double sum = 0.0; for (int k = 0; k < aCols; k++) { sum += a[i][k] * b[k][j]; } result[i][j] = sum; } } return result; } // 辅助方法:计算向量的外积 private static double[][] outerProduct(double[] a, double[] b) { int aLen = a.length; int bLen = b.length; double[][] result = new double[aLen][bLen]; for (int i = 0; i < aLen; i++) { for (int j = 0; j < bLen; j++) { result[i][j] = a[i] * b[j]; } } return result; } }
代码说明
- 参数调整:你可以修改
K值来改变分割的区域数量,比如想分成2个区域就设为2,4个就设为4 - 收敛条件:
THRESHOLD和MAX_ITERATIONS控制迭代停止的条件,可根据需求调整 - 可视化:代码里用簇的均值作为分割后区域的颜色,这样能更直观看到分割效果;如果你想保留原始像素颜色,只需要把
segmentedImage.setRGB(x, y, (r << 16) | (g << 8) | b);改成segmentedImage.setRGB(x, y, image.getRGB(x, y));即可 - 注意事项:如果图像太大,这个纯Java实现可能会有点慢,因为没有用到并行计算;另外初始化时随机选像素可能会导致每次分割结果略有不同,你也可以换成K-means初始化来提高稳定性
内容的提问来源于stack exchange,提问作者Walid Bousseta
相关产品推荐
相关产品推荐

