如何在LDA投影空间的旋转轴上绘制类投影数据直方图?
实现类均值连线平行方向的投影及上方直方图绘制
我用均值分离准则(寻找使类均值间距离最大化的参数w)和Fisher LDA给两个线性可分类别找了分隔线,画了直方图后发现重叠很多。现在想把数据点投影到与类均值连线平行的直线上,并且在这条直线上方绘制投影后的直方图(目标效果如图),但不知道怎么转成Python代码,求解决方法。

完整实现代码
import numpy as np import matplotlib.pyplot as plt from sklearn import datasets # 设置随机种子保证结果可复现 np.random.seed(8) X, y = datasets.make_blobs(n_samples=100, centers=2, n_features=2, center_box=(0, 10)) # 计算两类均值与均值分离方向向量w mu1 = np.mean(X[y == 0], axis=0) mu2 = np.mean(X[y == 1], axis=0) w = (mu2 - mu1) / np.linalg.norm(mu2 - mu1) # 计算数据点在w方向上的投影值和投影坐标 X_proj = np.dot(X, w) proj_points = X_proj[:, np.newaxis] * w # 投影到直线上的二维坐标 # 创建画布和主坐标轴 fig, ax = plt.subplots(figsize=(7, 7)) ax.set_xlim(0, 15) ax.set_ylim(0, 15) ax.set_xticks(np.arange(0, 15, 1)) ax.set_yticks(np.arange(0, 15, 1)) ax.grid(True) # 绘制原始数据点和类均值 ax.scatter(X[:, 0][y == 0], X[:, 1][y == 0], label='类别1', alpha=0.6) ax.scatter(X[:, 0][y == 1], X[:, 1][y == 1], label='类别2', alpha=0.6) ax.plot(mu1[0], mu1[1], 'X', color='red', markersize=10, label='类别1均值') ax.plot(mu2[0], mu2[1], 'X', color='red', markersize=10, label='类别2均值') # 绘制类均值连线和投影方向直线 ax.plot([mu1[0], mu2[0]], [mu1[1], mu2[1]], 'k--', label='类均值连线') # 延伸投影直线,覆盖所有投影点范围 proj_min = X_proj.min() - 2 proj_max = X_proj.max() + 2 line_start = proj_min * w line_end = proj_max * w ax.plot([line_start[0], line_end[0]], [line_start[1], line_end[1]], 'k-', label='投影直线') # 绘制原始点到投影点的连线 for x, p in zip(X, proj_points): ax.plot([x[0], p[0]], [x[1], p[1]], 'gray', linestyle=':', alpha=0.3) # 在投影直线上方绘制直方图 # 创建共享y轴的双x轴,用于放置直方图 hist_ax = ax.twiny() # 调整直方图位置,使其位于投影直线上方 hist_offset = 1.2 # 偏移量控制直方图与直线的距离 hist_y = line_start[1] + hist_offset + (line_end[1] - line_start[1]) * (X_proj - proj_min)/(proj_max - proj_min) # 绘制两类的直方图 hist_ax.hist(X_proj[y == 0], bins=8, alpha=0.5, label='类别1投影', orientation='horizontal') hist_ax.hist(X_proj[y == 1], bins=8, alpha=0.5, label='类别2投影', orientation='horizontal') # 调整直方图坐标轴,隐藏多余刻度 hist_ax.set_ylim(ax.get_ylim()) hist_ax.set_yticks([]) hist_ax.set_xlabel('投影值分布') # 添加图例 ax.legend(loc='upper left') hist_ax.legend(loc='upper right') plt.tight_layout() plt.show()
关键步骤说明
- 投影计算:通过
X_proj = np.dot(X, w)得到每个点在w方向上的投影值,再用proj_points = X_proj[:, np.newaxis] * w还原为二维平面上的投影坐标 - 投影直线绘制:根据投影值的极值延伸直线,确保覆盖所有数据点的投影范围
- 直方图放置:利用
twiny()创建双坐标轴,将直方图设置为水平方向,并通过偏移量调整其在投影直线上方的位置 - 投影连线:绘制原始点到投影点的灰色虚线,更直观展示投影过程
内容的提问来源于stack exchange,提问作者dg_m87
相关产品推荐
相关产品推荐

