1+ from __future__ import annotations
2+
13import os
24import re
5+ from dataclasses import dataclass
36from io import StringIO
7+ from typing import Mapping , Sequence
48
59import pandas as pd
610import requests
1620TABLE_NAME_PATTERN = re .compile (r"^[A-Za-z_][A-Za-z0-9_]*(\.[A-Za-z_][A-Za-z0-9_]*){2}$" )
1721
1822
23+ @dataclass (frozen = True )
24+ class TrainingResult :
25+ mse : float
26+ train_rows : int
27+ test_rows : int
28+ predicted_close_for_input : float
29+
30+ def to_dict (self ) -> dict [str , float | int ]:
31+ return {
32+ "mse" : self .mse ,
33+ "train_rows" : self .train_rows ,
34+ "test_rows" : self .test_rows ,
35+ "predicted_close_for_input" : self .predicted_close_for_input ,
36+ }
37+
38+
1939class DataProcessor :
2040 def __init__ (self , url : str ):
2141 self .url = url
@@ -47,7 +67,7 @@ def _validate_table_name(table_name: str) -> str:
4767 return table_name
4868
4969
50- def create_temp_table (conn , table_name : str ) -> None :
70+ def create_temp_table (conn : object , table_name : str ) -> None :
5171 table = _validate_table_name (table_name )
5272 with conn .cursor () as cur :
5373 cur .execute (f"DROP TABLE IF EXISTS { table } " )
@@ -57,7 +77,7 @@ def create_temp_table(conn, table_name: str) -> None:
5777 )
5878
5979
60- def insert_data_to_temp_table (conn , df : pd .DataFrame , table_name : str ) -> None :
80+ def insert_data_to_temp_table (conn : object , df : pd .DataFrame , table_name : str ) -> None :
6181 table = _validate_table_name (table_name )
6282 with conn .cursor () as cur :
6383 payload = df .copy ()
@@ -66,16 +86,53 @@ def insert_data_to_temp_table(conn, df: pd.DataFrame, table_name: str) -> None:
6686 cur .executemany (f"INSERT INTO { table } VALUES (%s, %s, %s)" , rows )
6787
6888
69- def fetch_data_from_temp_table (conn , table_name : str ) -> pd .DataFrame :
89+ def fetch_data_from_temp_table (conn : object , table_name : str ) -> pd .DataFrame :
7090 table = _validate_table_name (table_name )
7191 with conn .cursor () as cur :
7292 cur .execute (f"SELECT * FROM { table } " )
7393 rows = cur .fetchall ()
7494 return pd .DataFrame (rows , columns = ["DATE" , "CLOSE_PRICE" , "VOLATILITY_INDEX" ])
7595
7696
77- def run_training_and_generate_results () -> dict :
97+ def _fit_and_evaluate (
98+ features : pd .DataFrame , targets : pd .Series , predict_for : float
99+ ) -> TrainingResult :
100+ if features .empty or targets .empty :
101+ raise ValueError ("Training data is empty. Cannot fit model." )
102+
103+ X_train , X_test , y_train , y_test = train_test_split (
104+ features , targets , test_size = 0.2 , random_state = 42
105+ )
106+ if X_train .empty or X_test .empty :
107+ raise ValueError ("Insufficient rows after train-test split." )
108+
109+ model = LinearRegression ()
110+ model .fit (X_train , y_train )
111+
112+ predictions = model .predict (X_test )
113+ mse = mean_squared_error (y_test , predictions )
114+ single_prediction = float (model .predict (pd .DataFrame ({"VOLATILITY_INDEX" : [predict_for ]}))[0 ])
115+
116+ return TrainingResult (
117+ mse = float (mse ),
118+ train_rows = int (len (X_train )),
119+ test_rows = int (len (X_test )),
120+ predicted_close_for_input = single_prediction ,
121+ )
122+
123+
124+ def _build_report (result : TrainingResult ) -> str :
125+ return (
126+ "Our model predicts VIX CLOSE_PRICE from VOLATILITY_INDEX. "
127+ f"MSE on unseen data: { result .mse :.4f} . "
128+ f"Train rows: { result .train_rows } , Test rows: { result .test_rows } . "
129+ f"Sample prediction @ VOLATILITY_INDEX=0.40: { result .predicted_close_for_input :.2f} ."
130+ )
131+
132+
133+ def run_training_and_generate_results () -> Mapping [str , float | int ]:
78134 table_name = os .getenv ("SNOWFLAKE_TEMP_TABLE" , DEFAULT_TABLE )
135+ prediction_input = float (os .getenv ("VIX_PREDICTION_INPUT" , "0.40" ))
79136 conn = get_snowflake_connection ()
80137 try :
81138 data_processor = DataProcessor (VIX_DATA_URL )
@@ -86,23 +143,13 @@ def run_training_and_generate_results() -> dict:
86143 insert_data_to_temp_table (conn , df , table_name = table_name )
87144 fetched_df = fetch_data_from_temp_table (conn , table_name = table_name )
88145
89- X , y = fetched_df [["VOLATILITY_INDEX" ]], fetched_df ["CLOSE_PRICE" ]
90- X_train , X_test , y_train , y_test = train_test_split (
91- X , y , test_size = 0.2 , random_state = 42
92- )
93-
94- model = LinearRegression ()
95- model .fit (X_train , y_train )
96-
97- predictions = model .predict (X_test )
98- mse = mean_squared_error (y_test , predictions )
99-
100- report = (
101- "Our model predicts VIX CLOSE_PRICE from VOLATILITY_INDEX. "
102- f"The Mean Squared Error (MSE) on unseen test data is { mse :.4f} ."
146+ result = _fit_and_evaluate (
147+ features = fetched_df [["VOLATILITY_INDEX" ]],
148+ targets = fetched_df ["CLOSE_PRICE" ],
149+ predict_for = prediction_input ,
103150 )
104- print (report )
105- return { "mse" : float ( mse ), "train_rows" : int ( len ( X_train )), "test_rows" : int ( len ( X_test ))}
151+ print (_build_report ( result ) )
152+ return result . to_dict ()
106153 finally :
107154 conn .close ()
108155
0 commit comments