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

如何基于钻石数据集在散点图上绘制多类别多项式拟合模型

如何用diamonds数据集生成带多项式拟合的图表?

我正在使用标准diamonds数据集,需要创建类似下图的图表:
目标图表样式

目前我已实现两种绘图方式:

方式1:散点+折线图

import seaborn as sns
import matplotlib.pyplot as plt

# 加载数据
df = sns.load_dataset('diamonds')

plt.figure(figsize=(12, 8), dpi=200)

scatterplot = sns.scatterplot(data=df, x='carat', y='price', hue='cut', palette='viridis')

sns.lineplot(data=df, x='carat', y='price', hue='cut', palette='viridis', ax=scatterplot)

plt.xlabel('Carat')
plt.ylabel('Price')
plt.title('Scatter Plot of Price vs. Carat with Curved Lines (Viridis Palette)')

plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')

plt.show()

方式1生成的图表

方式2:分类别回归图

plt.figure(figsize=(12, 8), dpi=200)

cut_categories = df['cut'].unique()

for cut in cut_categories:
    data = df[df['cut'] == cut]
    sns.regplot(data=data, x='carat', y='price', scatter_kws={'s': 10}, label=cut)

plt.xlabel('Carat')
plt.ylabel('Price')
plt.title('Regression Plot of Price vs. Carat by Cut')

plt.legend(title='Cut')

plt.show()

方式2生成的图表

现在需要实现带有多项式拟合的图表,该怎么做?


实现多项式拟合的两种方法

方法1:用sns.regplot指定多项式阶数

sns.regplot本身支持多项式拟合,只需设置order参数指定多项式阶数(比如3阶),结合分类别循环即可实现:

import seaborn as sns
import matplotlib.pyplot as plt

df = sns.load_dataset('diamonds')

plt.figure(figsize=(12, 8), dpi=200)
palette = sns.color_palette('viridis', n_colors=len(df['cut'].unique()))

for idx, cut in enumerate(df['cut'].unique()):
    data = df[df['cut'] == cut]
    sns.regplot(
        data=data, x='carat', y='price',
        order=3,  # 指定多项式阶数,可根据数据趋势调整为2/3阶
        scatter_kws={'s': 10, 'alpha': 0.5},  # 降低散点透明度避免遮挡拟合线
        line_kws={'lw': 2},  # 加粗拟合线提升辨识度
        color=palette[idx],
        label=cut
    )

plt.xlabel('Carat')
plt.ylabel('Price')
plt.title('Price vs. Carat with Polynomial Fit (Order 3) by Cut')
plt.legend(title='Cut', bbox_to_anchor=(1.05, 1), loc='upper left')
plt.tight_layout()
plt.show()

方法2:用sns.lmplot批量生成多项式拟合图

如果想要更简洁的实现,sns.lmplot可以直接按hue分组生成多项式拟合,无需手动循环:

import seaborn as sns
import matplotlib.pyplot as plt

df = sns.load_dataset('diamonds')

g = sns.lmplot(
    data=df, x='carat', y='price',
    hue='cut', palette='viridis',
    order=3,  # 多项式阶数
    scatter_kws={'s': 10, 'alpha': 0.5},
    line_kws={'lw': 2},
    height=7, aspect=1.5,  # 控制图表大小比例
    facet_kws={'legend_out': True}
)

g.set_axis_labels('Carat', 'Price')
g.fig.suptitle('Price vs. Carat with Polynomial Fit (Order 3) by Cut', y=1.02)
plt.show()

关键说明

  • order参数:建议先尝试2或3阶,阶数过高容易出现过拟合,需根据数据实际趋势调整。
  • 散点透明度:设置alpha=0.5可以避免密集散点遮挡拟合线,提升图表可读性。
  • 两种方法各有优势:regplot循环方式更灵活控制样式,lmplot适合快速生成分组可视化。

内容的提问来源于stack exchange,提问作者Ghost Kid

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 10:25:22