Aliazimi00 commited on
Commit
e1812d5
·
verified ·
1 Parent(s): 89542d9

Update core/data.py

Browse files
Files changed (1) hide show
  1. core/data.py +19 -6
core/data.py CHANGED
@@ -8,19 +8,27 @@ try:
8
  except ImportError:
9
  raise ImportError("yfinance must be installed to fetch financial data.")
10
 
11
- def load_data(data_src="yahoo", ticker="AAPL", file_upload=None, start="2020-01-01", end="2023-01-01"):
 
 
 
12
  if data_src == "yahoo":
13
  try:
14
  info = yf.Ticker(ticker).info
15
  if not info:
16
  raise ValueError(f"Ticker '{ticker}' not found.")
17
- df = yf.download(ticker, start=start, end=end, progress=False)
 
 
18
  if df.empty:
19
  raise ValueError(f"No data found for ticker '{ticker}' in the specified date range. Please check the symbol and dates.")
 
 
 
 
 
20
  except Exception as e:
21
  raise ValueError(f"Error fetching data for ticker '{ticker}': {e}")
22
- df = df[['Close']].dropna().rename(columns={'Close': 'value'})
23
- df.reset_index(inplace=True)
24
  elif data_src == "csv":
25
  if file_upload is None:
26
  raise ValueError("CSV file upload required but not provided.")
@@ -33,10 +41,15 @@ def load_data(data_src="yahoo", ticker="AAPL", file_upload=None, start="2020-01-
33
  df = df[['Close']].rename(columns={'Close': 'value'})
34
  else:
35
  raise ValueError("CSV must contain a 'value' or 'Close' column.")
36
- df = df[['value']].dropna().reset_index(drop=True)
 
 
 
 
37
  else:
38
  raise ValueError("Invalid data source. 'csv' or 'yahoo' expected.")
39
- return df
 
40
 
41
 
42
  def preprocess_data(df, column, window_size=30):
 
8
  except ImportError:
9
  raise ImportError("yfinance must be installed to fetch financial data.")
10
 
11
+ def load_data(data_src="yahoo", ticker="AAPL", file_upload=None, start="2020-01-01", end="2023-01-01", horizon=1):
12
+ main_df = None
13
+ future_df = None
14
+
15
  if data_src == "yahoo":
16
  try:
17
  info = yf.Ticker(ticker).info
18
  if not info:
19
  raise ValueError(f"Ticker '{ticker}' not found.")
20
+ # Fetch data up to end_date + horizon for future actuals
21
+ extended_end = (pd.to_datetime(end) + pd.Timedelta(days=horizon)).strftime('%Y-%m-%d')
22
+ df = yf.download(ticker, start=start, end=extended_end, progress=False)
23
  if df.empty:
24
  raise ValueError(f"No data found for ticker '{ticker}' in the specified date range. Please check the symbol and dates.")
25
+ df = df[['Close']].dropna().rename(columns={'Close': 'value'})
26
+ df.reset_index(inplace=True)
27
+ # Split into main_df (up to end_date) and future_df (beyond end_date)
28
+ main_df = df[df['Date'] <= pd.to_datetime(end)].copy()
29
+ future_df = df[(df['Date'] > pd.to_datetime(end)) & (df['Date'] <= pd.to_datetime(extended_end))].copy()
30
  except Exception as e:
31
  raise ValueError(f"Error fetching data for ticker '{ticker}': {e}")
 
 
32
  elif data_src == "csv":
33
  if file_upload is None:
34
  raise ValueError("CSV file upload required but not provided.")
 
41
  df = df[['Close']].rename(columns={'Close': 'value'})
42
  else:
43
  raise ValueError("CSV must contain a 'value' or 'Close' column.")
44
+ df['Date'] = pd.to_datetime(df.get('Date', df.index))
45
+ df = df[['Date', 'value']].dropna().reset_index(drop=True)
46
+ # Split into main_df (up to end_date) and future_df (beyond end_date)
47
+ main_df = df[df['Date'] <= pd.to_datetime(end)].copy()
48
+ future_df = df[(df['Date'] > pd.to_datetime(end)) & (df['Date'] <= pd.to_datetime(end) + pd.Timedelta(days=horizon))].copy()
49
  else:
50
  raise ValueError("Invalid data source. 'csv' or 'yahoo' expected.")
51
+
52
+ return main_df, future_df
53
 
54
 
55
  def preprocess_data(df, column, window_size=30):