图表多线拟合及直线交点位置计算求助(附Python代码)
问题分析与解决方案
你的代码核心问题是每次都用全部数据集做拟合,导致三条线完全重合,根本没实现分段多线拟合;计算交点时也重复用同一组系数,结果毫无意义。要实现你想要的分段拟合+交点计算,得按以下步骤修改:
修正思路
- 先根据对数图的趋势,把数据分成多个区间(比如3段)
- 对每个区间的子数据单独做对数空间的线性拟合
- 计算相邻拟合直线的交点,再转换回原始坐标空间
完整可运行代码
import pandas as pd import numpy as np import matplotlib.pyplot as plt # 读取数据 dfn = pd.read_csv('input.csv') x = dfn['X'] y = dfn['Y'] # 转换为对数坐标 log_x = np.log(x) log_y = np.log(y) # -------------------------- # 第一步:手动指定分段阈值(根据你的图调整) # 这里示例用x的分位数,你可以换成自己需要的阈值 x_thresholds = np.quantile(x, [0.33, 0.66]) # 把数据分成3段 log_x_thresholds = np.log(x_thresholds) # -------------------------- # 初始化存储拟合系数的列表 coefs = [] # 绘制原始散点图 plt.scatter(log_x, log_y, label='原始数据', alpha=0.6) # 分段拟合并绘图 # 第一段:x <= 第一个阈值 mask1 = x <= x_thresholds[0] coef1 = np.polyfit(log_x[mask1], log_y[mask1], 1) poly1 = np.poly1d(coef1) plt.plot(log_x[mask1], poly1(log_x[mask1]), 'r--', label=f'分段1: y={np.exp(coef1[1]):.2f}*x^{coef1[0]:.2f}') coefs.append(coef1) # 第二段:第一个阈值 < x <= 第二个阈值 mask2 = (x > x_thresholds[0]) & (x <= x_thresholds[1]) coef2 = np.polyfit(log_x[mask2], log_y[mask2], 1) poly2 = np.poly1d(coef2) plt.plot(log_x[mask2], poly2(log_x[mask2]), 'g--', label=f'分段2: y={np.exp(coef2[1]):.2f}*x^{coef2[0]:.2f}') coefs.append(coef2) # 第三段:x > 第二个阈值 mask3 = x > x_thresholds[1] coef3 = np.polyfit(log_x[mask3], log_y[mask3], 1) poly3 = np.poly1d(coef3) plt.plot(log_x[mask3], poly3(log_x[mask3]), 'b--', label=f'分段3: y={np.exp(coef3[1]):.2f}*x^{coef3[0]:.2f}') coefs.append(coef3) # -------------------------- # 计算交点 intersection_points = [] for i in range(len(coefs)-1): # 取出相邻两条线的系数:log(y) = m1*log(x) + b1 ; log(y) = m2*log(x) + b2 m1, b1 = coefs[i] m2, b2 = coefs[i+1] # 求解交点的log(x):m1*logx + b1 = m2*logx + b2 if m1 == m2: print(f"第{i+1}和{i+2}条线平行,无交点") continue log_x_intersect = (b2 - b1) / (m1 - m2) log_y_intersect = m1 * log_x_intersect + b1 # 转换回原始坐标 x_intersect = np.exp(log_x_intersect) y_intersect = np.exp(log_y_intersect) intersection_points.append((x_intersect, y_intersect)) # 在图上标记交点 plt.scatter(log_x_intersect, log_y_intersect, color='k', s=50, zorder=10) # 保存交点数据 intersection_df = pd.DataFrame(intersection_points, columns=['X', 'Y']) intersection_df.to_csv('intersection_points.csv', index=False) print("交点数据已保存:") print(intersection_df) # 美化图表 plt.xlabel('log(X)') plt.ylabel('log(Y)') plt.legend() plt.grid(alpha=0.3) plt.show()
关键说明
- 分段阈值调整:示例用了x的分位数自动分段,你可以根据自己的图手动指定阈值(比如
x_thresholds = [10, 100]),确保每个分段的趋势符合线性 - 交点计算逻辑:在对数空间求解两条直线的交点,再通过
np.exp()转换回原始坐标 - 平行判断:如果两条线斜率相同,说明平行无交点,代码会跳过并提示
内容的提问来源于stack exchange,提问作者Amin Karimi
相关产品推荐
相关产品推荐

