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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 09:46:05