Update prep and predict functions
Change to new config dict
This commit is contained in:
1 parent
8878cbe305
commit
f3939c01dc
1 file changed
+39
-27
@@ -7,16 +7,14 @@ from sklearn.ensemble import RandomForestRegressor
|
|||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
|
|
||||||
from data.database import Base, dal
|
from data.database import Base, dal
|
||||||
import util.drawing as drawutil
|
from lottery_predictor.config import RANDOM_SEED, GAME_INFO
|
||||||
|
|
||||||
RANDOM_SEED = 42
|
|
||||||
|
|
||||||
|
|
||||||
def load_dataframe_by_dates(
|
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:
|
) -> 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 start_date: the start date for the data
|
||||||
:param end_date: the end date for the data
|
:param end_date: the end date for the data
|
||||||
|
|
||||||
@@ -26,17 +24,22 @@ def load_dataframe_by_dates(
|
|||||||
dal.connect()
|
dal.connect()
|
||||||
session = dal.Session()
|
session = dal.Session()
|
||||||
|
|
||||||
game_dict = drawutil.check_table_name(table_name=table_name)
|
game_dict = GAME_INFO[game]
|
||||||
if game_dict:
|
if game_dict:
|
||||||
target_table = Base.metadata.tables.get(game_dict['db_table_name'])
|
target_table = Base.metadata.tables.get(game_dict['db_table_name'])
|
||||||
if start_date is not None and end_date is not None:
|
if start_date is not None and end_date is not None:
|
||||||
sql_statement = (select(target_table)
|
sql_statement = (select(target_table)
|
||||||
.where(target_table.columns.draw_date >= start_date)
|
.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:
|
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:
|
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:
|
else:
|
||||||
sql_statement = select(target_table)
|
sql_statement = select(target_table)
|
||||||
|
|
||||||
@@ -45,9 +48,9 @@ def load_dataframe_by_dates(
|
|||||||
raise ValueError('An invalid table name was provided')
|
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)
|
:param limit: The N most recent draws (default: 10)
|
||||||
|
|
||||||
:returns: a pandas DataFrame containing the data
|
: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()
|
dal.connect()
|
||||||
session = dal.Session()
|
session = dal.Session()
|
||||||
|
|
||||||
game_dict = drawutil.check_table_name(table_name=table_name)
|
game_dict = GAME_INFO[game]
|
||||||
if game_dict:
|
if game_dict:
|
||||||
target_table = Base.metadata.tables.get(game_dict['db_table_name'])
|
target_table = Base.metadata.tables.get(game_dict['db_table_name'])
|
||||||
if limit is not None and limit > 0:
|
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]:
|
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
|
# 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
|
# Create X: flattened sliding windows
|
||||||
# Reshapes into (samples, window_size * columns)
|
x = np.array([clean_data[i : i + window_size].flatten() for i in indices])
|
||||||
X = np.array([data[i : i + window_size]. flatten() for i in indices])
|
|
||||||
|
|
||||||
# Create y_field: first 5 columns of the row which are the field balls
|
# Create y_field: first 5 columns (2D array)
|
||||||
y_field = data[window_size:, 0:5]
|
y_field = clean_data[window_size:, 0:5]
|
||||||
|
|
||||||
# Create y_game: 6th column (index 5) of the row which is the game ball
|
# Create y_game on the last column ensuring it is a 2D array
|
||||||
y_game = data[window_size:, 5]
|
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]:
|
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
|
data = data_frame.values
|
||||||
|
|
||||||
# Prepare and split the data
|
# 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
|
# Train the model for the field balls
|
||||||
field_model = RandomForestRegressor(n_estimators=200, random_state=RANDOM_SEED)
|
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
|
# Train the model for the game ball
|
||||||
game_model = RandomForestRegressor(n_estimators=200, random_state=RANDOM_SEED)
|
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
|
# 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
|
# Get predictions and round to the nearest whole number
|
||||||
pred_field = np.sort(np.round(field_model.predict(current_window)).astype(int))
|
predicted_field = np.sort(np.round(field_model.predict(current_window)).astype(int))
|
||||||
pred_game = np.round(game_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(
|
def get_most_common_number(
|
||||||
|
|||||||
Reference in new issue
Block a user