使用Julia Plots.jl绘制按行分组的矩阵数据遇问题求方案
Julia按分组变量z绘制散点图的正确实现
原代码的问题
size(y[:,1])返回的是(N,)元组,循环时i会取这个元组而非1到N的整数索引,直接导致索引越界错误- 给单个行的数据传入
group=z[:]逻辑错误,单组行数据只对应一个分组值,不需要传入整个z向量
正确实现方案
先按z的唯一值分组,提取每组对应的所有数据点再绘图,示例代码如下:
# 导入必要包 using Plots # 初始化空图 plt = scatter() # 获取所有唯一分组 unique_groups = unique(z) # 遍历每个分组绘图 for group in unique_groups # 找到当前分组对应的所有行索引 row_indices = findall(==(group), z) # 提取该分组的所有x、y数据并扁平化(将矩阵转为一维向量) x_points = vec(x[row_indices, :]) y_points = vec(y[row_indices, :]) # 添加当前分组的散点,用分组名作为标签 scatter!(plt, x_points, y_points, label=group) end # 显示图形 display(plt)
代码说明
- 先初始化空图
scatter(),避免循环中重复创建新图 findall(==(group), z)筛选出属于当前分组的所有行索引vec()将K×M的子矩阵转为一维向量,确保该分组的所有点都被正确绘制label=group用于区分不同分组的散点,替代原代码错误的group参数
如果需要更高效的分组操作,可以导入StatsBase包使用groupindices函数:
using Plots, StatsBase plt = scatter() # 获取唯一分组名和对应的索引范围 for (group_name, idx) in zip(unique(z), groupindices(z)) x_points = vec(x[idx, :]) y_points = vec(y[idx, :]) scatter!(plt, x_points, y_points, label=group_name) end display(plt)
内容的提问来源于stack exchange,提问作者Andrew
相关产品推荐
相关产品推荐

