如何基于本地CSV数据,在有线性趋势的时间区间内拟合回归直线?
针对时间序列线性回归的解决方案
嗨,我来帮你理清这个问题~首先明确回答你:你用本地CSV导入数据的情况,和你提到的那个问题处理逻辑基本一致——不管数据来源是网络还是本地文件,只要最终得到的是结构化的DataFrame(比如你代码里的amazon对象),后续拟合回归直线的步骤是完全通用的。
下面结合你的现有代码,给出具体的实现步骤:
1. 扩展现有代码,添加线性回归逻辑
首先,你需要导入线性回归相关的库,然后把时间转换成可用于回归的数值特征(因为日期类型不能直接作为模型输入),再拟合模型并绘制回归直线。
修改后的完整代码示例:
import datetime as dt import matplotlib.pyplot as plt import pandas as pd from sklearn.linear_model import LinearRegression # 新增线性回归库 import numpy as np plt.close('all') def parser(x): return pd.datetime.strptime(x, '%m/%d/%Y') # 读取本地CSV文件(你的原有代码) amazon = pd.read_csv('AMZN.csv', parse_dates=[1], index_col=1, squeeze=True, date_parser=parser) # ---------------------- 新增线性回归部分 ---------------------- # 将日期转换为数值特征:计算从数据起始日到当前日的天数 amazon['Days'] = (amazon.index - amazon.index[0]).days.values # 准备特征矩阵X和目标变量y X = amazon[['Days']] y = amazon['Close'] # 拟合线性回归模型 reg_model = LinearRegression() reg_model.fit(X, y) # 生成回归直线的预测值 amazon['Regression_Line'] = reg_model.predict(X) # ---------------------- 绘制图表 ---------------------- plt.figure(figsize=(12, 6)) plt.plot(amazon.index, amazon['Close'], label='AMZN Close Price', alpha=0.7) plt.plot(amazon.index, amazon['Regression_Line'], color='#ff4444', label='Linear Regression Line', linewidth=2) plt.title('AMZN Close Price with Linear Regression') plt.xlabel('Date') plt.ylabel('Price (USD)') plt.legend() plt.grid(alpha=0.3) plt.show()
2. 针对特定线性演化区间做回归
如果你只想在数据中某一段有明显线性趋势的时间区间拟合回归直线,可以先筛选出该区间的子集,再重复上述步骤:
# 示例:筛选2020年全年的数据(你可以替换成自己需要的区间) time_mask = (amazon.index >= '2020-01-01') & (amazon.index <= '2020-12-31') amazon_subset = amazon.loc[time_mask].copy() # 对子集拟合回归模型 X_sub = amazon_subset[['Days']] y_sub = amazon_subset['Close'] reg_model_sub = LinearRegression() reg_model_sub.fit(X_sub, y_sub) amazon_subset['Regression_Line'] = reg_model_sub.predict(X_sub) # 绘制子集的图表 plt.figure(figsize=(12, 6)) plt.plot(amazon_subset.index, amazon_subset['Close'], label='AMZN Close Price (2020)', alpha=0.7) plt.plot(amazon_subset.index, amazon_subset['Regression_Line'], color='#ff4444', label='Linear Regression Line', linewidth=2) plt.title('AMZN 2020 Close Price with Linear Regression') plt.xlabel('Date') plt.ylabel('Price (USD)') plt.legend() plt.grid(alpha=0.3) plt.show()
3. 为什么和你提到的问题处理方式一致?
那个问题中的数据来源是网络API,但最终得到的也是一个带日期索引的DataFrame。回归分析的核心是对结构化的数值数据进行拟合,和数据的获取方式无关——你用pd.read_csv导入本地CSV,和对方用API拉取数据,只是数据初始化的步骤不同,后续的回归建模、绘图逻辑完全相同。
内容的提问来源于stack exchange,提问作者Laura Marinescu
相关产品推荐
相关产品推荐

