散点图标记大小调整、标记类型自定义及警告排查问题
解决散点图大小失效、警告及标记类型自定义问题
嘿,我来帮你搞定这两个问题:
1. 标记大小不生效+RuntimeWarning的修复
你当前代码里犯了个小错误:同时用了s和sizes参数,这在plt.scatter里是冲突的!
s是直接给每个点指定具体大小数值,你的interval_size最大才2.4左右,直接用s=interval_size的话,点会小到几乎看不见,而且数值太小触发了matplotlib计算缩放时的警告。sizes是配合size参数用的,用来定义当你传入一个变量作为size时,大小的缩放范围。
修正后的代码如下,我还加了白边让小点也更清晰,可选添加大小图例:
plt.style.use("seaborn") sns.set_style("darkgrid") fig, ax = plt.subplots(figsize=(18,9)) # 核心修改:用size传入interval_size,sizes控制大小范围,确保变化明显 scatter = plt.scatter( x=true_labels, y=predictions, c=interval_size, alpha=0.6, cmap='viridis', size=interval_size, # 这里用size代替s sizes=(50, 300), # 调整这个范围让大小差异更直观 edgecolor='white' # 加白边提升小点的辨识度 ) cbar = plt.colorbar(scatter) cbar.set_label("Interval Width", labelpad=1, fontsize=20) # 可选:添加大小对应的图例 plt.legend(*scatter.legend_elements(prop="sizes", num=5), title="Interval Size") plt.title("True vs Predicted Labels", fontsize=36) plt.xlabel("True Labels", fontsize=25) plt.ylabel("Predicted Labels", fontsize=25) plt.show()
这样修改后,标记大小会随interval_size增大而明显变大,那个RuntimeWarning也会消失。
2. 按条件自定义标记类型
你提到的“按列设置不同标记类型”应该是表述有点偏差——毕竟每个点是(y, y_pred)的组合,没法直接按列分标记。我猜你是想根据某个条件(比如y和y_pred的大小关系,或者interval_size的阈值)来区分标记?
比如我们可以设置:当真实值y大于预测值y_pred时用圆形,反之用方形。实现方法是把数据拆成两组,分别调用plt.scatter:
plt.style.use("seaborn") sns.set_style("darkgrid") fig, ax = plt.subplots(figsize=(18,9)) # 定义区分条件:这里用y > y_pred为例 mask = true_labels > predictions # 第一组:y > y_pred,用圆形标记 scatter1 = ax.scatter( x=true_labels[mask], y=predictions[mask], c=interval_size[mask], alpha=0.6, cmap='viridis', size=interval_size[mask], sizes=(50, 300), marker='o', edgecolor='white' ) # 第二组:y <= y_pred,用方形标记 scatter2 = ax.scatter( x=true_labels[~mask], y=predictions[~mask], c=interval_size[~mask], alpha=0.6, cmap='viridis', size=interval_size[~mask], sizes=(50, 300), marker='s', edgecolor='white' ) # 设置颜色条和图例 cbar = plt.colorbar(scatter1) cbar.set_label("Interval Width", labelpad=1, fontsize=20) # 添加标记类型对应的图例 ax.legend([scatter1, scatter2], ['True > Predicted', 'True <= Predicted'], fontsize=15) plt.title("True vs Predicted Labels", fontsize=36) plt.xlabel("True Labels", fontsize=25) plt.ylabel("Predicted Labels", fontsize=25) plt.show()
如果你的区分逻辑是其他条件(比如interval_size是否大于1.5),只需要修改mask的表达式就行,比如mask = interval_size > 1.5。
内容的提问来源于stack exchange,提问作者bioinformatics_student
相关产品推荐
相关产品推荐

