如何在含断裂X轴与双断裂Y轴的子图中绘制跨图LOESS趋势线
带断裂轴的子图中绘制跨区域LOESS趋势线
问题背景
参考Matplotlib断裂Y轴的实现思路,基于合成数据搭建了包含断裂X轴与双断裂Y轴的子图布局:将多数常规数据放在下方主区域放大展示,异常值则置于上方的小区域。现需实现:使用sns.regplot绘制两条连续的局部加权(LOESS)趋势线,分别对应Cost与Hours字段,要求趋势线能跨越多张子图,适配Age>10的右侧区域以及上方异常值区域的维度。
合成数据
import pandas as pd import numpy as np import matplotlib.pyplot as plt import seaborn as sns from matplotlib.gridspec import GridSpec np.random.seed(123) # 生成合成数据:Age最小为1,Hours为正态分布,Cost为伽马分布缩放后的值 df = pd.DataFrame({ "Age": np.random.poisson(lam=7, size=1000) + 1, "Hours": np.random.normal(loc=2100, scale=100, size=1000), "Cost": np.random.gamma(shape=1.5, scale=1, size=1000) * 1000 })
当前实现代码(已完成布局与散点绘制)
figure = plt.figure(figsize=(10, 5)) # 划分2行2列的网格,设置高度、宽度比例 full_grid = GridSpec(nrows=2, ncols=2, height_ratios=[1, 4], width_ratios=[4, 1]) ax1 = figure.add_subplot(full_grid[0, 0]) # 左上:Age<=10的异常值区域 ax2 = figure.add_subplot(full_grid[0, 1]) # 右上:Age>10的异常值区域 ax3 = figure.add_subplot(full_grid[1, 0]) # 左下:Age<=10的常规数据区域 ax4 = figure.add_subplot(full_grid[1, 1]) # 右下:Age>10的常规数据区域 # 轴关联规则: # - ax1与ax3共享X轴范围;ax2与ax4共享X轴范围 # - ax3与ax4共享Y轴范围;ax1与ax2共享Y轴范围 # ---------------------- 下方主数据区域绘制 ---------------------- ax3hrs = ax3.twinx() # 用于绘制Hours的副轴 ax4hrs = ax4.twinx() # 绘制Age<=10的常规散点 sns.scatterplot(data=df[(df["Age"] <= 10) & (df["Cost"] <= df["Cost"].quantile(0.90))], x="Age", y="Cost", color="grey", alpha=0.1, ax=ax3) sns.scatterplot(data=df[(df["Age"] <= 10) & (df["Hours"] <= df["Hours"].quantile(0.90))], x="Age", y="Hours", color="grey", alpha=0.1, ax=ax3hrs) # 绘制Age>10的常规散点 sns.scatterplot(data=df[(df["Age"] > 10) & (df["Cost"] <= df["Cost"].quantile(0.90))], x="Age", y="Cost", color="grey", alpha=0.1, ax=ax4) sns.scatterplot(data=df[(df["Age"] > 10) & (df["Hours"] <= df["Hours"].quantile(0.90))], x="Age", y="Hours", color="grey", alpha=0.1, ax=ax4hrs) # 隐藏副轴的标签与刻度 ax3.set_ylabel(None); ax3.set_xlabel(None) ax3hrs.set_ylabel(None); ax3hrs.set_yticklabels("") ax3hrs.tick_params(axis="y", length=0) # 隐藏多余边框 ax3.spines[["right", "top"]].set_visible(False) ax3hrs.spines[["right", "top"]].set_visible(False) ax4.set_ylabel(None); ax4.set_xlabel(None); ax4.set_yticklabels("") ax4hrs.set_ylabel(None) ax4.tick_params(axis="y", length=0) ax4.spines[["left", "top"]].set_visible(False) ax4hrs.spines[["left", "top"]].set_visible(False) # 设置左下区域的X轴刻度 ax3.set_xticks(range(1, 11)) ax3.set_xticklabels(range(1, 11)) # 添加X轴断裂标记 d = 0.95 # 断裂斜线角度 break_kwargs = dict(marker=[(-1, -d), (1, d)], markersize=20, linestyle="none", color='k', mec='k', mew=1, clip_on=False) ax3.plot([1.015, 1.015], [0, 0.03], transform=ax3.transAxes, **break_kwargs) # ---------------------- 上方异常值区域绘制 ---------------------- ax1hrs = ax1.twinx() ax2hrs = ax2.twinx() # 同步轴范围 ax1.set_xlim(ax3.get_xlim()) ax2.set_xlim(ax4.get_xlim()) # 绘制Age<=10的异常值散点 sns.scatterplot(data=df[(df["Age"] <= 10) & (df["Cost"] > df["Cost"].quantile(0.90))], x="Age", y="Cost", color="grey", alpha=0.1, ax=ax1) sns.scatterplot(data=df[(df["Age"] <= 10) & (df["Hours"] > df["Hours"].quantile(0.90))], x="Age", y="Hours", color="grey", alpha=0.1, ax=ax1hrs) # 绘制Age>10的异常值散点 sns.scatterplot(data=df[(df["Age"] > 10) & (df["Cost"] > df["Cost"].quantile(0.90))], x="Age", y="Cost", color="grey", alpha=0.1, ax=ax2) sns.scatterplot(data=df[(df["Age"] > 10) & (df["Hours"] > df["Hours"].quantile(0.90))], x="Age", y="Hours", color="grey", alpha=0.1, ax=ax2hrs) # 隐藏多余标签与边框 ax1.set_xlabel(None); ax1.set_xticklabels(""); ax1.tick_params(axis="x", length=0) ax1.set_ylabel(None); ax1hrs.set_ylabel(None); ax1hrs.set_yticklabels("") ax1hrs.tick_params(axis="y", length=0) ax1.spines[["bottom", "right"]].set_visible(False) ax1hrs.spines[["bottom", "right"]].set_visible(False) ax2.set_xlabel(None); ax2.set_ylabel(None); ax2.set_xticklabels("") ax2.set_yticklabels(""); ax2.tick_params(axis="both", length=0) ax2hrs.set_xlabel(None); ax2hrs.set_ylabel(None) ax2.spines[["left", "bottom"]].set_visible(False) ax2hrs.spines[["left", "bottom"]].set_visible(False) # 添加Y轴断裂标记 ax3.plot([0.015, 0.015], [1.095, 1.05], transform=ax3.transAxes, **break_kwargs) ax4.plot([1.25, 1.25], [1.025, 1.075], transform=ax3.transAxes, **break_kwargs) plt.subplots_adjust(wspace=0.025)
实现跨区域LOESS趋势线
核心思路
先对全量数据拟合LOESS模型,生成完整的趋势线坐标点;再根据各子图的轴范围,拆分趋势线数据,将对应片段绘制到对应的子图中,保证线条在断裂处视觉连续。
完整代码(含趋势线)
import pandas as pd import numpy as np import matplotlib.pyplot as plt import seaborn as sns from matplotlib.gridspec import GridSpec np.random.seed(123) # 生成合成数据 df = pd.DataFrame({ "Age": np.random.poisson(lam=7, size=1000) + 1, "Hours": np.random.normal(loc=2100, scale=100, size=1000), "Cost": np.random.gamma(shape=1.5, scale=1, size=1000) * 1000 }) # 预计算分位数,用于划分常规值与异常值 cost_q90 = df["Cost"].quantile(0.90) hours_q90 = df["Hours"].quantile(0.90) figure = plt.figure(figsize=(10, 5)) full_grid = GridSpec(nrows=2, ncols=2, height_ratios=[1, 4], width_ratios=[4, 1]) ax1 = figure.add_subplot(full_grid[0, 0]) ax2 = figure.add_subplot(full_grid[0, 1]) ax3 = figure.add_subplot(full_grid[1, 0]) ax4 = figure.add_subplot(full_grid[1, 1]) # ---------------------- 拟合LOESS趋势线,生成全量趋势数据 ---------------------- # 拟合Cost的LOESS模型,获取趋势线数据 cost_reg = sns.regplot(data=df, x="Age", y="Cost", lowess=True, scatter=False, ax=plt.gca()) cost_line = cost_reg.get_lines()[0] x_cost, y_cost = cost_line.get_data() plt.close() # 关闭临时绘图 # 拟合Hours的LOESS模型,获取趋势线数据 hours_reg = sns.regplot(data=df, x="Age", y="Hours", lowess=True, scatter=False, ax=plt.gca()) hours_line = hours_reg.get_lines()[0] x_hours, y_hours = hours_line.get_data() plt.close() # ---------------------- 下方主数据区域 ---------------------- ax3hrs = ax3.twinx() ax4hrs = ax4.twinx() # 绘制散点 sns.scatterplot(data=df[(df["Age"] <=10) & (df["Cost"] <= cost_q90)], x="Age", y="Cost", color="grey", alpha=0.1, ax=ax3) sns.scatterplot(data=df[(df["Age"] <=10) & (df["Hours"] <= hours_q90)], x="Age", y="Hours", color="grey", alpha=0.1, ax=ax3hrs) sns.scatterplot(data=df[(df["Age"] >10) & (df["Cost"] <= cost_q90)], x="Age", y="Cost", color="grey", alpha=0.1, ax=ax4) sns.scatterplot(data=df[(df["Age"] >10) & (df["Hours"] <= hours_q90)], x="Age", y="Hours", color="grey", alpha=0.1, ax=ax4hrs) # 绘制Cost趋势线的常规值片段 mask_cost_regular = (x_cost <=10) & (y_cost <= cost_q90) ax3.plot(x_cost[mask_cost_regular], y_cost[mask_cost_regular], color="darkorange", linewidth=2) mask_cost_regular_right = (x_cost >10) & (y_cost <= cost_q90) ax4.plot(x_cost[mask_cost_regular_right], y_cost[mask_cost_regular_right], color="darkorange", linewidth=2) # 绘制Hours趋势线的常规值片段 mask_hours_regular = (x_hours <=10) & (y_hours <= hours_q90) ax3hrs.plot(x_hours[mask_hours_regular], y_hours[mask_hours_regular], color="darkblue", linewidth=2) mask_hours_regular_right = (x_hours >10) & (y_hours <= hours_q90) ax4hrs.plot(x_hours[mask_hours_regular_right], y_hours[mask_hours_regular_right], color="darkblue", linewidth=2) # 轴样式调整 ax3.set_ylabel(None); ax3.set_xlabel(None) ax3hrs.set_ylabel(None); ax3hrs.set_yticklabels(""); ax3hrs.tick_params(axis="y", length=0) ax3.spines[["right", "top"]].set_visible(False); ax3hrs.spines[["right", "top"]].set_visible(False) ax4.set_ylabel(None); ax4.set_xlabel(None); ax4.set_yticklabels("") ax4hrs.set_ylabel(None); ax4.tick_params(axis="y", length=0) ax4.spines[["left", "top"]].set_visible(False); ax4hrs.spines[["left", "top"]].set_visible(False) ax3.set_xticks(range(1,11)); ax3.set_xticklabels(range(1,11)) # 添加X轴断裂标记 d = 0.95 break_kwargs = dict(marker=[(-1, -d), (1, d)], markersize=20, linestyle="none", color='k', mec='k', mew=1, clip_on=False) ax3.plot([1.015,1.015], [0,0.03], transform=ax3.transAxes, **break_kwargs) # ---------------------- 上方异常值区域 ---------------------- ax1hrs = ax1.twinx() ax2hrs = ax2.twinx() ax1.set_xlim(ax3.get_xlim()); ax2.set_xlim(ax4.get_xlim()) # 绘制散点 sns.scatterplot(data=df[(df["Age"] <=10) & (df["Cost"] > cost_q90)], x="Age", y="Cost", color="grey", alpha=0.1, ax=ax1) sns.scatterplot(data=df[(df["Age"] <=10) & (df["Hours"] > hours_q90)], x="Age", y="Hours", color="grey", alpha=0.1, ax=ax1hrs) sns.scatterplot(data=df[(df["Age"] >10) & (df["Cost"] > cost_q90)], x="Age", y="Cost", color="grey", alpha=0.1, ax=ax2) sns.scatterplot(data=df[(df["Age"] >10) & (df["Hours"] > hours_q90)], x="Age", y="Hours", color="grey", alpha=0.1, ax=ax2hrs) # 绘制Cost趋势线的异常值片段 mask_cost_outlier = (x_cost <=10) & (y_cost > cost_q90) ax1.plot(x_cost[mask_cost_outlier], y_cost[mask_cost_outlier], color="darkorange", linewidth=2) mask_cost_outlier_right = (x_cost >10) & (y_cost > cost_q90) ax2.plot(x_cost[mask_cost_outlier_right], y_cost[mask_cost_outlier_right], color="darkorange", linewidth=2) # 绘制Hours趋势线的异常值片段 mask_hours_outlier = (x_hours <=10) & (y_hours > hours_q90) ax1hrs.plot(x_hours[mask_hours_outlier], y_hours[mask_hours_outlier], color="darkblue", linewidth=2) mask_hours_outlier_right = (x_hours >10) & (y_hours > hours_q90) ax2hrs.plot(x_hours[mask_hours_outlier_right], y_hours[mask_hours_outlier_right], color="darkblue", linewidth=2) # 轴样式调整 ax1.set_xlabel(None); ax1.set_xticklabels(""); ax1.tick_params(axis="x", length=0) ax1.set_ylabel(None); ax1hrs.set_ylabel(None); ax1hrs.set_yticklabels(""); ax1hrs.tick_params(axis="y", length=0) ax1.spines[["bottom", "right"]].set_visible(False); ax1hrs.spines[["bottom", "right"]].set_visible(False) ax2.set_xlabel(None); ax2.set_ylabel(None); ax2.set_xticklabels("") ax2.set_yticklabels(""); ax2.tick_params(axis="both", length=0) ax2hrs.set_xlabel(None); ax2hrs.set_ylabel(None) ax2.spines[["left", "bottom"]].set_visible(False); ax2hrs.spines[["left", "bottom"]].set_visible(False) # 添加Y轴断裂标记 ax3.plot([0.015,0.015], [1.
相关产品推荐
相关产品推荐

