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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 10:38:32