From f3939c01dc0eb5b2d2e9041661951187171ec4ba Mon Sep 17 00:00:00 2001 From: Chris Smith Date: Tue, 26 May 2026 19:28:48 -0400 Subject: [PATCH] Update prep and predict functions Change to new config dict --- lottery_predictor/analyze.py | 66 +++++++++++++++++++++--------------- 1 file changed, 39 insertions(+), 27 deletions(-) diff --git a/lottery_predictor/analyze.py b/lottery_predictor/analyze.py index 773a4a8..e884881 100644 --- a/lottery_predictor/analyze.py +++ b/lottery_predictor/analyze.py @@ -7,16 +7,14 @@ from sklearn.ensemble import RandomForestRegressor from sqlalchemy import select from data.database import Base, dal -import util.drawing as drawutil - -RANDOM_SEED = 42 +from lottery_predictor.config import RANDOM_SEED, GAME_INFO def load_dataframe_by_dates( - table_name: str, start_date: datetime | None = None, end_date: datetime | None = None + game: str, start_date: datetime | None = None, end_date: datetime | None = None ) -> pd.DataFrame: """ - :param table_name: the name of the table to load data from + :param game: the name of the game to get data for :param start_date: the start date for the data :param end_date: the end date for the data @@ -26,17 +24,22 @@ def load_dataframe_by_dates( dal.connect() session = dal.Session() - game_dict = drawutil.check_table_name(table_name=table_name) + game_dict = GAME_INFO[game] if game_dict: target_table = Base.metadata.tables.get(game_dict['db_table_name']) if start_date is not None and end_date is not None: sql_statement = (select(target_table) .where(target_table.columns.draw_date >= start_date) - .where(target_table.columns.draw_date <= end_date)) + .where(target_table.columns.draw_date <= end_date) + .order_by(target_table.columns.draw_date.desc())) elif start_date is not None and end_date is None: - sql_statement = (select(target_table).where(target_table.columns.draw_date >= start_date)) + sql_statement = (select(target_table) + .where(target_table.columns.draw_date >= start_date) + .order_by(target_table.columns.draw_date.desc())) elif start_date is None and end_date is not None: - sql_statement = (select(target_table).where(target_table.columns.draw_date <= end_date)) + sql_statement = (select(target_table) + .where(target_table.columns.draw_date <= end_date) + .order_by(target_table.columns.draw_date.desc())) else: sql_statement = select(target_table) @@ -45,9 +48,9 @@ def load_dataframe_by_dates( raise ValueError('An invalid table name was provided') -def load_dataframe_most_recent(table_name: str, limit: int = 10) -> pd.DataFrame: +def load_dataframe_most_recent(game: str, limit: int = 10) -> pd.DataFrame: """ - :param table_name: the name of the table to load data from + :param game: the name of the game to get data for :param limit: The N most recent draws (default: 10) :returns: a pandas DataFrame containing the data @@ -56,7 +59,7 @@ def load_dataframe_most_recent(table_name: str, limit: int = 10) -> pd.DataFrame dal.connect() session = dal.Session() - game_dict = drawutil.check_table_name(table_name=table_name) + game_dict = GAME_INFO[game] if game_dict: target_table = Base.metadata.tables.get(game_dict['db_table_name']) if limit is not None and limit > 0: @@ -71,21 +74,29 @@ def load_dataframe_most_recent(table_name: str, limit: int = 10) -> pd.DataFrame def prepare_split_data(data: np.ndarray, window_size: int = 10) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + # Clean the data by removing the draw_date and multiplier columns + clean_data = data[:, 1:7] + + # Check to be sure there is enough data for the window_size + if len(clean_data) <= window_size: + raise ValueError( + f"Not enough data! Dataset has {len(clean_data)} rows, " + f"but window_size requires at least {window_size + 1} rows." + ) # Calculate the indices for all windows at once - indices = np.arange(len(data) - window_size) + indices = np.arange(len(clean_data) - window_size) # Create X: flattened sliding windows - # Reshapes into (samples, window_size * columns) - X = np.array([data[i : i + window_size]. flatten() for i in indices]) + x = np.array([clean_data[i : i + window_size].flatten() for i in indices]) - # Create y_field: first 5 columns of the row which are the field balls - y_field = data[window_size:, 0:5] + # Create y_field: first 5 columns (2D array) + y_field = clean_data[window_size:, 0:5] - # Create y_game: 6th column (index 5) of the row which is the game ball - y_game = data[window_size:, 5] + # Create y_game on the last column ensuring it is a 2D array + y_game = clean_data[window_size:, 5].ravel() - return X, y_field, y_game + return x, y_field, y_game def make_prediction(data_frame: pd.DataFrame, window_size: int = 10) -> tuple[np.ndarray, np.ndarray]: @@ -93,23 +104,24 @@ def make_prediction(data_frame: pd.DataFrame, window_size: int = 10) -> tuple[np data = data_frame.values # Prepare and split the data - X, y_field, y_game = prepare_split_data(data, window_size=window_size) + x, y_field, y_game = prepare_split_data(data, window_size=window_size) # Train the model for the field balls field_model = RandomForestRegressor(n_estimators=200, random_state=RANDOM_SEED) - field_model.fit(X, y_field) + field_model.fit(x, y_field) # Train the model for the game ball game_model = RandomForestRegressor(n_estimators=200, random_state=RANDOM_SEED) - game_model.fit(X, y_game) + game_model.fit(x, y_game) # Predict the next draw - current_window = data[-window_size:].flatten().reshape(1, -1) + clean_data = data[:, 1:7] + current_window = clean_data[-window_size:].flatten().reshape(1, -1) # Get predictions and round to the nearest whole number - pred_field = np.sort(np.round(field_model.predict(current_window)).astype(int)) - pred_game = np.round(game_model.predict(current_window)).astype(int) + predicted_field = np.sort(np.round(field_model.predict(current_window)).astype(int)) + predicted_game = np.round(game_model.predict(current_window)).astype(int) - return pred_field[0], pred_game[0] + return predicted_field[0], predicted_game[0] def get_most_common_number(