-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
64 lines (54 loc) · 2.3 KB
/
Copy pathmain.py
File metadata and controls
64 lines (54 loc) · 2.3 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
import pandas as pd
import argparse
from models.arima_model import arima_forecast
from models.holt_winters_model import holt_winters_forecast
from models.prophet_model import prophet_forecast
from models.tbats_model import tbats_forecast
from models.ensemble_model import ensemble_forecast
from utils.preprocessing import preprocess_series
def main():
i = 3
parser = argparse.ArgumentParser()
parser.add_argument('--input', type=str, required=True, help='Path to input CSV')
parser.add_argument('--horizon', type=int, required=True, help='Forecast horizon')
parser.add_argument('--models', type=str, default='arima,holt,prophet,tbats,ensemble',
help='Comma-separated list of models to run')
args = parser.parse_args()
df = pd.read_csv(args.input)
series = preprocess_series(df)
model_list = [m.strip().lower() for m in args.models.split(',')]
forecasts = {}
trained_models = {}
for model_name in model_list:
print(f"Running model: {model_name}")
if model_name == 'arima':
forecast, model = arima_forecast(series, args.horizon)
forecasts['arima'] = forecast
trained_models['arima'] = model
elif model_name == 'holt':
forecast, model = holt_winters_forecast(series, args.horizon)
forecasts['holt'] = forecast
trained_models['holt'] = model
elif model_name == 'prophet':
forecast, model = prophet_forecast(series, args.horizon)
forecasts['prophet'] = forecast
trained_models['prophet'] = model
elif model_name == 'tbats':
forecast, model = tbats_forecast(series, args.horizon)
forecasts['tbats'] = forecast
trained_models['tbats'] = model
elif model_name == 'ensemble':
continue
else:
print(f"⚠️ Unknown model: {model_name}")
for name, forecast in forecasts.items():
forecast.index = range(1, len(forecast) + 1)
if 'ensemble' in model_list:
print("Running ensemble model")
ensemble = ensemble_forecast(forecasts)
forecasts['ensemble'] = ensemble
forecast_df = pd.DataFrame(forecasts)
forecast_df.to_csv(f'forecast{i}.csv', index=False)
print("✅ Forecast saved to forecast.csv")
if __name__ == "__main__":
main()