Matplotlib:单X值对应多Y值的散点图及回归线绘制问题
问题解答:Matplotlib散点图的数据类型与指数关系回归线绘制
一、单X值对应多Y值的最优数据类型选择
针对你这种一个X对应多个Y值的算法性能评估场景,我推荐两种实用的数据类型,具体选哪个看你的后续需求:
- 嵌套列表(直观易用):最直接的方式,把每个X对应的Y值打包成子列表,比如:
绘图时用循环遍历每个X和对应的Y组,就能快速画出散点:x = [1, 2, 3, 4, 5] y_groups = [[1, 2, 3], [4, 5, 6], [7, 8, 9], [10, 11, 12], [13, 14, 15]]import matplotlib.pyplot as plt for xi, yi in zip(x, y_groups): plt.scatter([xi]*len(yi), yi, label=f'x={xi}') plt.legend() plt.show() - Pandas DataFrame(推荐用于性能评估):如果需要后续做统计分析(比如计算每个X对应的Y均值、方差,这对评估算法稳定性很重要),把数据整理成长格式的DataFrame会更方便:
绘图时直接调用import pandas as pd # 构造长格式数据 data = [] for xi, yi in zip(x, y_groups): for y_val in yi: data.append({'x': xi, 'y': y_val}) df = pd.DataFrame(data)scatter即可,后续还能轻松计算统计量并可视化:
这种格式能让你更清晰地观察算法在不同X下的性能分布,非常适合你的评估需求。plt.scatter(df['x'], df['y']) # 添加每个X对应的均值线,辅助评估性能集中趋势 mean_y = df.groupby('x')['y'].mean() plt.plot(mean_y.index, mean_y.values, color='red', marker='o', label='Mean Y') plt.legend() plt.show()
二、指数关系下绘制回归线的方法
当然可以绘制回归线!不过因为你的X和Y是指数关系(比如Y = a*b^X),直接拟合线性回归线会偏差很大,推荐两种靠谱的方法:
方法1:对数变换转线性拟合
把指数关系转化为线性关系:对Y取自然对数,得到ln(Y) = ln(a) + X*ln(b),这就变成了标准线性模型,拟合后再转换回指数形式即可:
import numpy as np from sklearn.linear_model import LinearRegression # 把数据展开成一维数组 x_flat = np.repeat(x, 3).reshape(-1, 1) # 每个X重复3次,对应3个Y值 y_flat = np.concatenate(y_groups) # 对数变换 ln_y = np.log(y_flat) # 拟合线性模型 model = LinearRegression() model.fit(x_flat, ln_y) # 预测并转换回指数形式 x_pred = np.linspace(min(x), max(x), 100).reshape(-1, 1) ln_y_pred = model.predict(x_pred) y_pred = np.exp(ln_y_pred) # 绘图 plt.scatter(x_flat, y_flat, label='Raw Data') plt.plot(x_pred, y_pred, color='green', linewidth=2, label='Exponential Regression') plt.legend() plt.xlabel('X') plt.ylabel('Y') plt.show()
方法2:直接非线性拟合(更贴合原始数据)
用scipy.optimize.curve_fit直接拟合指数函数,不需要做变换,结果会更准确:
from scipy.optimize import curve_fit # 定义指数函数模型 def exp_func(x, a, b): return a * (b ** x) # 拟合最优参数 params, _ = curve_fit(exp_func, x_flat.flatten(), y_flat) a_opt, b_opt = params # 生成预测曲线 y_pred = exp_func(x_pred.flatten(), a_opt, b_opt) # 绘图 plt.scatter(x_flat, y_flat, label='Raw Data') plt.plot(x_pred, y_pred, color='orange', linewidth=2, label='Non-linear Exponential Fit') plt.legend() plt.xlabel('X') plt.ylabel('Y') plt.show()
这种方法不需要假设变换后的线性关系,拟合结果完全贴合原始数据的指数趋势。
内容的提问来源于stack exchange,提问作者Bartholomas
相关产品推荐
相关产品推荐

