import datetime import logging import os import warnings import numpy as np import pandas as pd from pandas.errors import PerformanceWarning from tadas.domain.features import apply_adj_to_dod, add_is_any, combine_data_sources from tadas.domain.trending_flags import apply_trending_flags from tadas.platform import context as contexts from tadas.snowflake import ( days_trending_generator, facts_source, model_registry, snowflake_publish, tables, ) from tadas.platform import config from tadas.platform import dynamodb as dynamodb_utils from tadas.snowflake import client as snowflake_utils logger = logging.getLogger(__name__) warnings.simplefilter(action='ignore', category=PerformanceWarning) def inference(report, report_date: datetime.date, model_config): if not _need_to_run(report, report_date, model_config): logger.info("No need to run model") return start_time = datetime.datetime.now() input_table = _collect_data(report, report_date, model_config) output_table = tables.get_table_name('output', model_config.MODEL_VERSION, report) model_registry.model_predict_proba( input_table=input_table, output_table=output_table, ml_model_name=model_config.SNOWFLAKE_REGISTRY_MODEL_NAME, ml_model_version=model_config.SNOWFLAKE_REGISTRY_MODEL_VERSION, model_features=model_config.ML_MODEL_FEATURES, ) snowflake_publish.sync_historical( source_table=output_table, target_table=tables.TADAS_HISTORICAL_TABLE, model_version=model_config.MODEL_VERSION, ) dynamodb_utils.set_overall_status( feed_name=contexts.get_feed_name(report=report), date=str(report_date), overall_status=dynamodb_utils.STATUS_DOWNLOADED, ) delta = datetime.datetime.now() - start_time logger.info(f'Total Time: {delta.total_seconds() / 60.0:.1f} minutes') logger.info("Model inference completed.") def _need_to_run(report, report_date, model_config): output_table = tables.get_table_name('output', model_config.MODEL_VERSION, report) return not facts_source.is_table_fresh( table_name=output_table, expected_date=report_date, ) def _collect_data(report, report_date, model_config): tracks_df = _create_tracks_table(report_date, report, model_config) logger.info(f"{datetime.datetime.now().strftime('%H:%M:%S')} : starting generate data") data_df = _generate_data(tracks_df, report_date, report, model_config) days_trending_df = days_trending_generator.version_v3( report_date=report_date, group_by_country=report != 'global', select_tracks_table=tables.get_table_name('select_tracks', model_config.MODEL_VERSION, report), ) if config.get('SAVE_DAYS_TRENDING'): days_trending_generator.save_days_trending( df=days_trending_df, table_name=tables.get_table_name('days_trending', model_config.MODEL_VERSION, report), ) combined_df = combine_data_sources(data_df, days_trending_df, report_date) apply_trending_flags(combined_df, model_config.TRENDING_FLAGS) combined_table = tables.get_table_name('combined', model_config.MODEL_VERSION, report) snowflake_publish.save_combined_df(combined_df, table_name=combined_table) return combined_table def _create_tracks_table(report_date, report, model_config): tracks_df = facts_source.select_tracks( report_date, group_by_country=report != 'global', ) tracks_table_name = tables.get_table_name('select_tracks', model_config.MODEL_VERSION, report) track_df_clone = tracks_df.copy() track_df_clone['report_date'] = pd.to_datetime(report_date).date() snowflake_utils.saveto_snowflake( df=track_df_clone, myschema=os.getenv('SNOWFLAKE_SCHEMA'), table=tracks_table_name, mode='replace', ) return tracks_df def _generate_data(tracks_df, yesterday, report, model_config): output_df = tracks_df.copy() logger.info(f'Expected rows: {len(output_df)}') output_df = _get_all_streams(output_df, yesterday, report, model_config) output_df = apply_adj_to_dod(output_df) output_df = add_is_any(output_df) output_df = output_df.fillna(0.0) # Sometimes we get inf growth, so we're capping it at 1000% output_df = output_df.replace([np.inf, -np.inf], 1000.0) output_df['pfn_geo'] = output_df['isrc_cd'].astype(str) + '_' + output_df['geo_country'].astype(str) return output_df def _get_all_streams(output_df, current_date, report, model_config): tracks_table = tables.get_table_name('select_tracks', model_config.MODEL_VERSION, report) with snowflake_utils.snowflake_connection() as conn: for source_table, selected_columns in model_config.DBT_MODELS_COLUMNS.items(): logger.info(f'Loading {len(selected_columns)} columns from {source_table}') query = facts_source.prepare_sql_select_from_dbt_model( source_table=source_table, tracks_table=tracks_table, columns_definition=selected_columns, group_by_country=report == 'countries', periods=model_config.PERIODS, ) params = {'current_date': current_date} s_df = snowflake_utils.query_snowflake_to_df(connection=conn, query=query, params=params) logger.info(f'Loaded {len(s_df)} rows from {source_table}') for column_name in selected_columns.keys(): _apply_wow_and_dod(s_df, column_name, model_config) s_df.columns = [x.lower() for x in s_df.columns] output_df = pd.merge(output_df, s_df, how='left', on=['isrc_cd', 'geo_country']) logger.info(f'Merged to output_df. Total size is {len(output_df)} rows.') return output_df def _apply_wow_and_dod(s_df, column_name, model_config): yest_x = model_config.PREFIX_YEST + column_name seven_x = model_config.PREFIX_SEVEN + column_name wow_new_col = 'wow_' + column_name dod_new_col = 'dod_' + column_name s_df[wow_new_col] = s_df[column_name] / s_df[seven_x].replace(0, np.nan) - 1 s_df[dod_new_col] = s_df[column_name] / s_df[yest_x].replace(0, np.nan) - 1 # Some columns come back as Decimal which can't be converted by pyarrow. Force float. s_df[column_name] = s_df[column_name].apply(float) s_df[yest_x] = s_df[yest_x].apply(float) s_df[seven_x] = s_df[seven_x].apply(float) s_df[wow_new_col] = s_df[wow_new_col].apply(float) s_df[dod_new_col] = s_df[dod_new_col].apply(float)