-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathrun_pipeline.py
More file actions
345 lines (271 loc) · 11 KB
/
Copy pathrun_pipeline.py
File metadata and controls
345 lines (271 loc) · 11 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
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
# Copyright (c) 2025 Soares
#
# SPDX-License-Identifier: Apache-2.0
"""
PriceSentinel: Event-Aware Energy Price Forecasting
Main command-line interface for running the forecasting pipeline.
Usage:
# Run full pipeline
python run_pipeline.py --country PT --all --start-date 2023-01-01 --end-date 2024-12-31
# Run individual stages
python run_pipeline.py --country PT --fetch --start-date 2024-01-01 --end-date 2024-01-31
python run_pipeline.py --country PT --clean
python run_pipeline.py --country PT --train
python run_pipeline.py --country PT --forecast
# Get pipeline info
python run_pipeline.py --country PT --info
""" # noqa: E501
from __future__ import annotations
import argparse
import asyncio
import sys
from datetime import datetime
from pathlib import Path
# Add the project root to the path
sys.path.insert(0, str(Path(__file__).parent))
from config.country_registry import CountryRegistry
from core.logging_config import setup_logging
from core.pipeline import Pipeline
from core.pipeline_builder import PipelineBuilder
from data_fetchers import auto_register_countries
def parse_arguments() -> argparse.Namespace:
"""Parse command-line arguments."""
parser = argparse.ArgumentParser(
description="PriceSentinel: Event-Aware Energy Price Forecasting",
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""
Examples:
# Run full pipeline for Portugal
python run_pipeline.py --country PT --all --start-date 2023-01-01 --end-date 2024-12-31
# Fetch data only
python run_pipeline.py --country PT --fetch --start-date 2024-01-01 --end-date 2024-01-31
# Run forecast only
python run_pipeline.py --country PT --forecast --forecast-date 2025-01-07
# Get information about available data
python run_pipeline.py --country PT --info
""", # noqa: E501
)
# Required arguments
parser.add_argument(
"--country", type=str, required=True, help="Country code (e.g., PT, ES, DE, XX for mock)"
)
# Pipeline stages
parser.add_argument("--all", action="store_true", help="Run all pipeline stages")
parser.add_argument("--fetch", action="store_true", help="Fetch raw data from APIs")
parser.add_argument("--clean", action="store_true", help="Clean and verify data")
parser.add_argument("--features", action="store_true", help="Engineer features")
parser.add_argument("--train", action="store_true", help="Train forecasting model")
parser.add_argument("--forecast", action="store_true", help="Generate forecasts")
parser.add_argument("--info", action="store_true", help="Show pipeline and data information")
parser.add_argument(
"--model-name",
type=str,
default="baseline",
help="Model name to use for training (default: baseline)",
)
# Date arguments
parser.add_argument("--start-date", type=str, help="Start date (YYYY-MM-DD)")
parser.add_argument("--end-date", type=str, help="End date (YYYY-MM-DD)")
parser.add_argument(
"--forecast-date", type=str, help="Forecast date (YYYY-MM-DD), defaults to today"
)
parser.add_argument(
"--forecast-start-date",
type=str,
help="Start date for forecast range (YYYY-MM-DD)",
)
parser.add_argument(
"--forecast-end-date",
type=str,
help="End date for forecast range (YYYY-MM-DD)",
)
# Logging
parser.add_argument(
"--log-level",
type=str,
choices=["DEBUG", "INFO", "WARNING", "ERROR"],
default="INFO",
help="Logging level (default: INFO)",
)
parser.add_argument(
"--fast-train",
action="store_true",
help="Use fast training mode (e.g., smaller model or shorter run)",
)
return parser.parse_args()
def validate_arguments(args: argparse.Namespace) -> bool:
"""
Validate command-line arguments.
Args:
args: Parsed arguments
Returns:
True if valid, False otherwise
"""
# Check if at least one action is specified
actions = [
args.all,
args.fetch,
args.clean,
args.features,
args.train,
args.forecast,
args.info,
]
if not any(actions):
print("Error: No action specified. Use --all or specify individual stages.")
print("Use --help for usage information.")
return False
# Check date requirements
if args.all or args.fetch:
if not args.start_date or not args.end_date:
print("Error: --start-date and --end-date are required for fetching data.")
return False
# Validate date format
try:
datetime.strptime(args.start_date, "%Y-%m-%d")
datetime.strptime(args.end_date, "%Y-%m-%d")
except ValueError:
print("Error: Dates must be in YYYY-MM-DD format.")
return False
# Check forecast date / range requirements for forecast-only runs
if args.forecast and not args.all:
has_single = bool(args.forecast_date)
has_range = bool(args.forecast_start_date and args.forecast_end_date)
if not (has_single or has_range):
print(
"Error: Provide either --forecast-date or "
"--forecast-start-date and --forecast-end-date for forecasting."
)
return False
if has_range:
try:
start_dt = datetime.strptime(args.forecast_start_date, "%Y-%m-%d")
end_dt = datetime.strptime(args.forecast_end_date, "%Y-%m-%d")
except ValueError:
print("Error: Forecast dates must be in YYYY-MM-DD format.")
return False
if start_dt > end_dt:
print("Error: --forecast-start-date must be <= --forecast-end-date.")
return False
return True
class PipelineCLI:
"""CLI runner for the PriceSentinel pipeline."""
def __init__(self, args: argparse.Namespace) -> None:
import logging
self.args = args
# Initialise logger immediately to avoid Optional type issues
self.logger: logging.Logger = logging.getLogger(__name__)
# Pipeline is initialised later in init_pipeline
self.pipeline: Pipeline | None = None
def setup_logging(self) -> None:
setup_logging(level=self.args.log_level)
@staticmethod
def print_header() -> None:
print("\n" + "=" * 70)
print("PriceSentinel: Event-Aware Energy Price Forecasting")
print("=" * 70 + "\n")
def register_countries(self) -> list[str]:
self.logger.info("Registering countries...")
auto_register_countries()
available_countries = CountryRegistry.list_countries()
self.logger.info(f"Available countries: {', '.join(available_countries)}")
return available_countries
def validate_country(self, available_countries: list[str]) -> None:
if self.args.country.upper() not in available_countries:
self.logger.error(
f"Country '{self.args.country}' not registered.\n"
f"Available countries: {available_countries}"
)
sys.exit(1)
def init_pipeline(self) -> None:
self.logger.info(f"Initializing pipeline for {self.args.country}...")
self.pipeline = PipelineBuilder.create_pipeline(self.args.country)
def show_info_and_exit(self) -> None:
if self.pipeline is None:
raise RuntimeError("Pipeline not initialized")
info = self.pipeline.get_info()
print("\n" + "=" * 70)
print(f"PIPELINE INFORMATION: {info['country_code']}")
print("=" * 70)
print(f"Country: {info['country_name']}")
print(f"Timezone: {info['timezone']}")
print(f"Run ID: {info['run_id']}")
print(f"Data directory: {info['data_directory']}")
print("\nData info:")
for key, value in info["data_info"].items():
if key != "sources":
print(f" {key}: {value}")
if "sources" in info["data_info"]:
print(" Sources:")
for source, count in info["data_info"]["sources"].items():
print(f" - {source}: {count} files")
print("=" * 70 + "\n")
sys.exit(0)
async def run_stages(self) -> None:
if self.pipeline is None:
raise RuntimeError("Pipeline not initialized")
# Determine effective model name, incorporating fast-train flag if set
model_name = self.args.model_name
if getattr(self.args, "fast_train", False) and not model_name.endswith("_fast"):
model_name = f"{model_name}_fast"
# Date range (may be None; pipeline methods fall back to last fetched range)
start_date = self.args.start_date
end_date = self.args.end_date
if self.args.all:
await self.pipeline.run_full_pipeline(
self.args.start_date,
self.args.end_date,
self.args.forecast_date,
model_name=model_name,
)
return
if self.args.fetch:
await self.pipeline.fetch_data(start_date, end_date)
if self.args.clean:
self.pipeline.clean_and_verify(start_date, end_date)
if self.args.features:
self.pipeline.engineer_features(start_date, end_date)
if self.args.train:
self.pipeline.train_model(start_date, end_date, model_name=model_name)
if self.args.forecast:
# Use range if provided, otherwise single forecast date
if self.args.forecast_start_date and self.args.forecast_end_date:
self.pipeline.generate_forecast_range(
self.args.forecast_start_date,
self.args.forecast_end_date,
model_name=model_name,
)
else:
self.pipeline.generate_forecast(self.args.forecast_date, model_name=model_name)
def run(self) -> None:
self.setup_logging()
self.print_header()
available_countries = self.register_countries()
self.validate_country(available_countries)
try:
self.init_pipeline()
if self.args.info:
self.show_info_and_exit()
asyncio.run(self.run_stages())
print("\n" + "=" * 70)
print(f"[OK] Pipeline completed successfully for {self.args.country}")
print("=" * 70 + "\n")
except KeyboardInterrupt:
self.logger.warning("\n\nPipeline interrupted by user")
sys.exit(1)
except Exception as e:
self.logger.error(f"\n\nPipeline failed: {e}", exc_info=True)
print("\n" + "=" * 70)
print(f"[FAIL] Pipeline failed for {self.args.country}")
print(f"Error: {e}")
print("=" * 70 + "\n")
sys.exit(1)
def main() -> None:
"""Main entry point for the CLI."""
args = parse_arguments()
if not validate_arguments(args):
sys.exit(1)
cli = PipelineCLI(args)
cli.run()
if __name__ == "__main__":
main()