简介:这套MATLAB程序聚焦高斯混合模型GMM与高斯混合回归GMR两大核心算法,主要面向机器学习初学者、相关课程学员以及需要快速上手概率建模的工程师。程序以清晰脚本分步实现GMM的初始化E步M步迭代与参数估计,并进一步扩展GMR完成连续变量回归;压缩包中既包含基础演示demo,也涵盖任务参数化张量GMM、LQR控制等进阶示例,配合运行效果图与数据文件,有助于将理论公式与各类应用场景打通。压缩包文件总数26个,包含19个m文件(核心函数、演示脚本及工具箱)、6个mat数据文件及1个txt说明文档,整体仅98KB,轻量易用。目前已有1958人学习下载,适合通过动手调参和运行demo深入理解GMM聚类与GMR预测的细节,也为后续研究任务参数化或动态系统建模提供了可参考的实现模板。 搞过一段时间机器人轨迹学习和运动规划的朋友,应该都绕不开“高斯混合模型GMM”和“高斯混合回归GMR”这两个名字。当时我做人体运动数据建模,需要从一堆带有明显多模态特征的位置点里提取规律,再用提取出的模型去预测新输入下的输出,GMM+GMR这套组合几乎是按图索骥的最佳选择。配合MATLAB自带的统计工具箱,从拟合高斯混合分布到推导条件回归结果,一条龙下来比想象中顺手,但坑也不少。这篇博文就围绕GMM和GMM+GMR(高斯混合回归)的实现细节展开,给出可直接跑通的MATLAB程序结构,也会把我实际踩过的初始化、收敛、协方差退化这几个典型的坑一并说来。
这套内容适合对概率模型有点概念、但想快速在MATLAB里落地实验的读者,也适合做控制、机器人、信号处理方向需要一种轻量非线性回归方法的朋友参考。不搞数学推导堆砌,重点放在怎么调、怎么用、怎么避坑。
1. 内容整体设计与思路拆解
1.1 为什么是“GMM+GMR”而不是直接回归
传统非线性回归,比如多项式拟合、BP神经网络、支持向量机回归,都能做曲线拟合。但它们有一个共性:建模的是一个确定性映射,最多再给个全局噪声方差。这在实际数据里往往不够。以人体运动轨迹为例,同一个起点到达同一个终点,可能有多条合理的路径,数据分布会出现明显的“多峰”特征。此时用单一函数去拟合,结果就是各个路径的“平均态”,既不符合物理直觉,也不会被下游控制模块接受。
GMM先把数据分布看成若干高斯分量的叠加,每个分量捕获一个局部簇结构。这样天然尊重了数据的多模态特性。而GMR则利用GMM联合分布里的输入-输出相关性,对输入做条件期望计算,得到一条连续、光滑的回归曲线。整个过程不需要人为指定基函数,也不用提前知道分簇数。只要分量数适中、初始化合理,模型自己就能找出数据里的结构。
从工程上讲,这套方案还有一个好处:训练与推理逻辑统一。训练时用的是EM(期望最大化)算法,推理时只需把输入代入条件高斯公式,不需要像神经网络那样额外撸一段反向传播代码。
1.2 程序整体架构与各模块职责
我在MATLAB里的实现分成四个模块,职责边界清楚,后面调试和维护都方便:
- 数据生成/导入模块:负责生成仿真数据或者读入外部数据。仿真数据建议用几个已知均值和协方差的高斯分布混在一起生成,这样能随时对照真实参数验证拟合效果。
- GMM训练模块:核心调用
fitgmdist函数,负责指定分量数K、协方差类型、正则化系数和最大迭代次数。 - GMR推理模块:手动实现“已知输入 -> 求每个分量的条件均值和条件协方差 -> 加权融合”的流程。这是整个程序里最需要理解清楚的部分,不能偷懒直接调工具箱。
- 可视化与评估模块:画原始数据散点图、拟合的椭圆等高线、回归曲线和±1σ置信带。
模块之间用函数封装,避免把一堆代码堆在一个脚本里。主程序只负责设置参数和调用。
2. 核心细节解析与实操要点
2.1 fitgmdist的核心参数与选型逻辑
MATLAB里拟合GMM的最直接方式是调用fitgmdist,我常用的一组参数结构如下:
gm = fitgmdist(X, K, ... 'Start', 'plus', ... 'CovarianceType', 'full', ... 'RegularizationValue', 0.01, ... 'Options', statset('MaxIter', 1000, 'TolFun', 1e-6));这里有几个参数值得逐一说明,因为它们直接影响拟合效果和稳定性。
- K(分量数):这是GMM里最敏感的超参数。K太小,模型欠拟合,多模态信息丢失;K太大,过拟合,某些分量会缩成只包裹了一两个样本点的“细针”。确定K我的经验是先用
fitgmdist跑K从1到10的循环,记录每个K下的AIC或BIC,选一个拐点处的K。AIC容易选多,BIC容易选少,实际任务里我一般倾向BIC往小一点选,保证模型泛化能力。 - CovarianceType:'full'允许每个分量有独立的全协方差矩阵,能捕捉变量之间的旋转相关关系。'diagonal'假设变量间独立,参数少、训练快,但表达力不足。在GMR场景里,我们要利用输入和输出分块协方差之间的关系,所以必须用'full',否则回归会退化成多个独立的一维高斯映射。
- RegularizationValue:这个参数是给协方差矩阵对角线加一个很小的正数,防止矩阵奇异。很多人忽略它,导致跑高维数据时报
Ill-conditioned covariance错误。我通常设0.001到0.1之间,具体看数据量级。 - Start:EM算法是迭代求解,初始点选不好容易陷入局部最优。'plus'方法会先对数据做k-means++聚类,然后用聚类结果构造初始参数,比随机初始化稳定很多。数据量大时也可以传一个包含K个样本点的初始均值矩阵,加速收敛。
2.2 EM算法本质与实现中的“隐变量”直觉
EM算法在GMM里的本质是:因为我们不知道每个样本来自哪个高斯分量,所以就把“样本归属”当作隐变量。E步固定当前参数,计算每个样本属于每个分量的后验概率(responsibility);M步用这些后验概率做加权统计,重新估计均值、协方差和权重。两步交替迭代,直到对数似然变化小于阈值。
我的理解方式比较朴素:这就像老师带一群学生分成K个学习小组。E步是老师根据目前各小组的水平差异,给每个学生重新分配小组的概率;M步是根据新的分组概率,重新评估每个小组的平均水平和波动范围。来回几次之后,分组和小组特征都趋于稳定。
在MATLAB里,fitgmdist已经封装好这些细节,不需要自己写EM迭代。但理解这个交替过程对调试很有帮助。比如你发现拟合结果每次跑都不一样,那就是EM收敛到了不同局部最优,这时应该检查初始化方式,而不是怀疑代码写错。
2.3 数据预处理不能忽略的规范问题
GMM对数据的尺度极其敏感。如果一个特征的量纲是0到1,另一个是0到10000,协方差矩阵会被大尺度特征主导,小尺度特征的信息基本会被淹没。我一般先做z-score标准化:
mu_data = mean(X); std_data = std(X); X_norm = (X - mu_data) ./ std_data;训练完模型后,GMR预测时也要把输入先按同样的参数标准化,输出再反转回去。这个细节很多人会忘,导致预测结果直接飞掉。
还有一点:如果输入数据有NaN,fitgmdist会直接报错或者静默丢弃,导致输出维度对不上。我在处理外部导入的数据时,会显式检查一下:
assert(~any(isnan(X(:))), '数据包含NaN,请先处理缺失值');3. 实操过程与核心环节实现
3.1 仿真数据构造与训练脚本
先给一段完整的仿真数据脚本。我们用三个高斯分布混合生成数据,其中一个维度的均值有明显分段结构,模拟一个“输入-输出”关系带有多峰特征的任务。
rng(42); % 三个高斯分量:均值、协方差、样本数 mu1 = [2, 5]; Sigma1 = [0.8, 0.3; 0.3, 0.6]; n1 = 300; mu2 = [6, 8]; Sigma2 = [1.2, -0.4; -0.4, 0.7]; n2 = 300; mu3 = [9, 3]; Sigma3 = [0.6, 0.1; 0.1, 0.5]; n3 = 300; X1 = mvnrnd(mu1, Sigma1, n1); X2 = mvnrnd(mu2, Sigma2, n2); X3 = mvnrnd(mu3, Sigma3, n3); X = [X1; X2; X3];这里我故意把三个分量的均值在x方向上错开(2、6、9),并且在y方向上有重叠。这样就构造出了一个随着x增大,输出y不再是一条简单直线的数据结构,GMR在这种数据上能展示出它比线性回归和多项式回归更强表达能力的特点。
训练部分:
K = 3; gm = fitgmdist(X, K, ... 'Start', 'plus', ... 'CovarianceType', 'full', ... 'RegularizationValue', 0.01, ... 'Options', statset('MaxIter', 1000, 'TolFun', 1e-6));3个真分量,K=3可以直接还原。实际操作中不知道真实分量数,我会跑一个BIC曲线选择脚本再定,后面专门讲。
3.2 GMR推导与实现:从联合高斯到条件高斯
GMR的核心公式不复杂,但必须要理解它的推导脉络。我们把每个样本看成输入x和输出y拼接而成的向量[x; y],假设它服从一个高斯混合分布。对于第k个分量,有:
均值: mu_k = [mu_xk; mu_yk] 协方差: Sigma_k = [Sigma_xxk, Sigma_xyk; Sigma_yxk, Sigma_yyk]已知输入为x_query时,在第k个分量下的输出y的条件分布仍然是高斯分布:
mu_y_given_x_k = mu_yk + Sigma_yxk * inv(Sigma_xxk) * (x_query - mu_xk) Sigma_y_given_x_k = Sigma_yyk - Sigma_yxk * inv(Sigma_xxk) * Sigma_xyk这里的关键在于:每个分量给出的条件均值是线性的,但不同分量的线性系数不同,最后加权混合后,整体映射就变成了分段近似线性、总体呈非线性的函数。这就是GMR表达力来源。
多分量融合时还要算每个分量对当前查询点的归一化权重。这个权重不是简单用训练时的混合系数,而是“混合系数乘以该输入点在第k个高斯分量下的概率密度”再归一化:
w_k = alpha_k * mvnpdf(x_query, mu_xk, Sigma_xxk) / sum_j alpha_j * mvnpdf(x_query, mu_xj, Sigma_xxj)最终输出条件期望为:
E[y|x_query] = sum_k w_k * mu_y_given_x_k条件协方差为:
Var[y|x_query] = sum_k w_k^2 * Sigma_y_given_x_k第二项我习惯取平方加权,得到的置信带更保守,比直接线性加权更容易画出好看的区间。
对应的MATLAB实现函数如下:
function [y_pred, y_var] = gmr_predict(gm, x_query, dim_in) % dim_in: 输入维度,输入数据必须是前dim_in列 K = gm.NumComponents; d = size(gm.mu, 2); dim_out = d - dim_in; alpha = gm.ComponentProportion; mu_all = gm.mu; Sigma_all = gm.Sigma; n_query = size(x_query, 1); y_pred = zeros(n_query, dim_out); y_var = zeros(n_query, dim_out); for i = 1:n_query xq = x_query(i, :); w = zeros(K, 1); mu_cond = zeros(K, dim_out); Sigma_cond = zeros(K, dim_out, dim_out); for k = 1:K mu_k = mu_all(k, :); Sigma_k = Sigma_all(:, :, k); mu_xk = mu_k(1:dim_in); mu_yk = mu_k(dim_in+1:end); Sigma_xx = Sigma_k(1:dim_in, 1:dim_in); Sigma_xy = Sigma_k(1:dim_in, dim_in+1:end); Sigma_yy = Sigma_k(dim_in+1:end, dim_in+1:end); % 条件均值与协方差 inv_Sxx = Sigma_xx \ eye(dim_in); mu_cond(k, :) = mu_yk + (xq - mu_xk) * inv_Sxx * Sigma_xy; Sigma_cond(k, :, :) = Sigma_yy - Sigma_xy' * inv_Sxx * Sigma_xy; % 未归一化权重 w(k) = alpha(k) * mvnpdf(xq, mu_xk, Sigma_xx); end w = w / sum(w); y_pred(i, :) = w' * mu_cond; for k = 1:K y_var(i, :) = y_var(i, :) + w(k)^2 * squeeze(Sigma_cond(k, :, :))'; end end end这段代码里用Sigma_xx \ eye(dim_in)求逆,比直接inv(Sigma_xx)稳定一些,尤其是在协方差矩阵条件数较大时,矩阵左除在数值上更可靠。数据量特别大时可以用pagefun或者mldivide做批量加速,但一般教学和科研场景下循环就够了。
有一点要提醒:fitgmdist输出的Sigma是一个d x d x K的三维数组,不是元胞数组,所以取第k个分量的协方差矩阵要用Sigma_all(:,:,k)。我见过不少人在这里踩坑,直接用Sigma_all(k)取,报错找半天原因。
3.3 完整回归流程与可视化
算完预测值后,把结果画出来,我一般同时画三类图:
- 原始数据散点图:画前两列数据,直观看到聚类结构。
- GMM拟合结果图:用
gscatter按后验概率最大分量着色,再用ezplot或fcontour画每个分量的1σ椭圆。这一步可以快速判断K选得对不对。 - GMR回归曲线图:在x方向等间隔取100个查询点,调用
gmr_predict,画均值曲线和置信带。
x_query = linspace(min(X(:,1))-0.5, max(X(:,1))+0.5, 100)'; x_query = [x_query, zeros(100, 1)]; % 第二列占位,实际会被条件化掉 [y_pred, y_var] = gmr_predict(gm, x_query, 1); std_band = sqrt(y_var); figure; hold on; scatter(X(:,1), X(:,2), 10, [0.6 0.6 0.6], 'filled'); plot(x_query(:,1), y_pred, 'r-', 'LineWidth', 2); plot(x_query(:,1), y_pred + 1.96*std_band, 'r--', 'LineWidth', 1.2); plot(x_query(:,1), y_pred - 1.96*std_band, 'r--', 'LineWidth', 1.2); xlabel('x'); ylabel('y'); legend('数据点','GMR预测','95%置信带');这里x_query第二列填0,是因为我们给gmr_predict传入的是完整特征向量,程序会按dim_in=1自动忽略后一维。实际使用中如果输入输出维度已经拆开,直接传单列x_query就行。
4. 高斯混合回归的进阶用法与影响范围
4.1 多输入多输出场景的扩展
GMR并不局限于单输入单输出。把dim_in改成2或者更大,输入输出按拼接顺序约定好,代码几乎不用改。比如在机器人运动规划里,经常把时间t和当前位置x作为输入,预测下一个位置dx,这就是一个高维回归。GMR的好处是它自动处理输入输出之间的相关结构,不需要为每个输出单独建一个回归模型。
我在实验里试过输入为2维、输出为3维的GMR,效果依然稳定。关键在于协方差矩阵的分块要对应清楚,千万别把输入顺序和输出顺序搞混,否则条件公式里的Sigma_xy取错块,预测结果完全没意义。
4.2 与神经网络、高斯过程回归的对比
很多入门者会纠结:既然有神经网络和高斯过程回归,为什么还要用GMR?
- 神经网络强在特征自动提取和海量数据,但需要调的结构参数多,训练时间长,可解释性偏差。
- 高斯过程回归(GPR)理论上比你优秀,小样本拟合效果极好,还自带不确定性估计,但复杂度是O(n³),数据超过几千个点就开始卡顿。
- GMR恰好站在两者中间。它有显式的概率结构,可以解释每个分量意味着什么模式;训练复杂度比GPR低得多;预测时是线性运算,速度快;还能天然处理多模态分布,这一步GPR是做不到的。
在数据量中等(几百到几千样本)、维度不高(三维到十维)、又需要概率输出和可解释性的场景下,GMR的优势很明显。如果数据量再大、复杂度再高,就考虑深度生成模型或者变分方法了。
5. 常见问题与排查技巧实录
5.1 分量数K该怎么选
K的选择没有标准答案,我的经验是固定随机种子跑BIC曲线,选拐点。示例代码如下:
rng(42); K_list = 1:8; BIC = zeros(size(K_list)); for i = 1:length(K_list) gm_temp = fitgmdist(X, K_list(i), ... 'Start', 'plus', ... 'CovarianceType', 'full', ... 'RegularizationValue', 0.01); BIC(i) = gm_temp.BIC; end [~, best_idx] = min(BIC); K_opt = K_list(best_idx);用BIC而不是AIC,是为了在模型复杂度和拟合度之间取平衡。BIC惩罚项更重,得到的K偏小,模型更简洁。对回归任务来说,稍微偏小的K通常比稍微偏大的K更容易得到一个光滑、稳定的映射函数。
还有个实用技巧:跑完BIC后人工看一眼聚类结果。如果某个分量的标准差在某一维上趋近于0,说明K偏大了,这个分量在硬拟合离群点。此时手动减K,比纯看BIC拐点更可靠。
5.2 初始值导致的结果不稳定
我遇到过一模一样的代码和数据,连续跑两次,结果完全不同。原因就是EM收敛到了不同的局部最优解。解决方式有几个层次:
- 最直接的是固定随机种子,保证实验可复现。
- 更治本的是用
Start参数的'plus'选项,先做k-means++得到更合理的初始化。 - 如果还不行,手动指定
'Start'为一个K x d的矩阵,选取数据里分散的K个点作为初始均值。这个方法在数据规模不大时效果显著。
顺带说一句,MATLAB 2020版本之后fitgmdist里Start选项支持传struct,可以分别指定初始均值、协方差和权重。追求极致稳定的话,可以先用k-means聚类结果构造初始参数,再丢给EM迭代。
5.3 协方差矩阵奇异与数值溢出问题
协方差矩阵奇异是GMM实现里的高频问题,典型症状是MATLAB警告“Ill-conditioned covariance”或者程序直接中断。出现条件主要有三种:
- 数据本身存在共线性(两个维度几乎完全相关)。
- K设得太大,某个分量只剩下少数几个点,协方差估计退化。
- 数据里存在极端离群点,拉大了协方差的尺度。
我的排查顺序是:先检查数据相关性矩阵,看是否存在近线性相关的维度;再把K往下调,观察是否还报错;最后给RegularizationValue调大一点,从0.01试到0.1。一般前两步能解决90%的问题。
数值溢出主要出现在mvnpdf计算时,当查询点距离某个分量中心非常远,密度值会下溢为0,导致后续权重归一化出现除零错误。我的处理方式是给密度值统一加一个小正数防止极端情况,或者改用log-sum-exp技巧,先取对数再归一化,数值稳定性好很多。不过日常实验加个小常数就够了。
写作时的几点体会
这套GMM+GMR程序写下来,我最深的感触是:数学公式看着抽象,一旦落到矩阵操作里,反而变得很具体。条件高斯分布的推导可以写满两页纸,但MATLAB里实现就是四五个矩阵乘积的事。关键是分清楚哪个矩阵是Sigma_xx、哪个是Sigma_xy、哪个是Sigma_yy,分块对了,一切水到渠成。
如果后续想扩展,可以试试把GMM换成变分贝叶斯GMM,让模型自动决定分量数,省去BIC这一层。或者把GMR的协方差输出接进一个轨迹规划器里,用不确定性信息引导搜索更安全的路径。这些都是很有意思的方向,有机会再单独开一篇聊。
本文还有配套的精品资源,点击获取