CoolFace
Apppublic

datajoi/Dataset-Test-Workflow

sourceHugging Facemitupdated 2y agoView on Hugging Face
3likes
app.py380 linesDownload Raw Back to root
1import os2import json3import duckdb4import gradio as gr5import pandas as pd6import pandera as pa7from pandera import Column8import ydata_profiling as pp9from langchain_huggingface import HuggingFaceEndpoint, ChatHuggingFace10from langsmith import traceable11from langchain import hub12import warnings13import dlt14warnings.filterwarnings("ignore", category=DeprecationWarning)15 16# Height of the Tabs Text Area17TAB_LINES = 818 19 20#----------CONNECT TO DATABASE----------21md_token = os.getenv('MD_TOKEN')22conn = duckdb.connect(f"md:my_db?motherduck_token={md_token}", read_only=True)23#---------------------------------------24 25 26#-------LOAD HUGGINGFACE-------27models = ["Qwen/Qwen2.5-72B-Instruct","meta-llama/Meta-Llama-3-70B-Instruct",28          "meta-llama/Llama-3.1-70B-Instruct"]29 30model_loaded = False 31for model in models:32  try:33      endpoint = HuggingFaceEndpoint(repo_id=model, max_new_tokens=8192)34      info = endpoint.client.get_endpoint_info()35      model_loaded = True36      break37  except Exception as e:38      print(f"Error for model {model}: {e}")39      continue40  41llm = ChatHuggingFace(llm=endpoint).bind(max_tokens=8192)42#---------------------------------------43 44#-----LOAD PROMPT FROM LANCHAIN HUB-----45prompt_autogenerate = hub.pull("autogenerate-rules-testworkflow")46prompt_user_input = hub.pull("usergenerate-rules-testworkflow")47 48#--------------ALL UTILS----------------49# Get Databases50def get_schemas():51    schemas = conn.execute("""52    SELECT DISTINCT schema_name53    FROM information_schema.schemata54    WHERE schema_name NOT IN ('information_schema', 'pg_catalog')55    """).fetchall()56    return [item[0] for item in schemas]57 58# Get Tables59def get_tables_names(schema_name):60    tables = conn.execute(f"SELECT table_name FROM information_schema.tables WHERE table_schema = '{schema_name}'").fetchall()61    return [table[0] for table in tables]62 63# Update Tables64def update_table_names(schema_name):65    tables = get_tables_names(schema_name)66    return gr.update(choices=tables)67# def get_data_df(schema):68#     print('Getting Dataframe from the Database')69#     return conn.sql(f"SELECT * FROM {schema} LIMIT 1000")70 71@dlt.resource72def fetch_data(schema):73    result = conn.sql(f"SELECT * FROM {schema} LIMIT 1000")74    75    while True:76        chunk_df = result.fetch_df_chunk(2)77        78        if chunk_df is None or len(chunk_df) == 0:79            break80        else:81            yield chunk_df82 83def create_pipeline(schema):84    dataset_name = schema.split('.')[1]85    print("Dataset Name: ", dataset_name)86    87    table_name = schema.split('.')[2]88    print("Table Name: ", table_name)89    90    pipeline =dlt.pipeline(91        pipeline_name='duckdb_pipeline',92        destination='duckdb',93        dataset_name= dataset_name,94    )95    96    load_info = pipeline.run(fetch_data(schema), table_name = table_name,97                             write_disposition = "replace")98    99    print(load_info)100    return dataset_name + "." + table_name 101 102def load_pipeline(table_name):103    _conn = duckdb.connect("duckdb_pipeline.duckdb")104    return _conn, _conn.sql(f"SELECT * FROM {table_name} LIMIT 1000").df()105 106def df_summary(df):107    summary = []108 109    for column in df.columns:110        if pd.api.types.is_numeric_dtype(df[column]):111            summary.append({112                "column": column,113                "max": df[column].max(),114                "min": df[column].min(),115                "count": df[column].count(),116                "nunique": df[column].nunique(),117                "dtype": str(df[column].dtype),118                "top": None119            })120            121        elif pd.api.types.is_categorical_dtype(df[column]) or pd.api.types.is_object_dtype(df[column]):122            top_value = df[column].mode().iloc[0] if not df[column].mode().empty else None123            124            summary.append({125                "column": column,126                "max": None,  127                "min": None, 128                "count": df[column].count(),129                "nunique": df[column].nunique(),130                "dtype": str(df[column].dtype),131                "top": top_value132            })133    summary_df = pd.DataFrame(summary)134    return summary_df.reset_index(drop=True)135 136def format_prompt(df):137    summary = df_summary(df)138    return prompt_autogenerate.format_prompt(data=df.head().to_json(orient='records'),139                                           summary=summary.to_json(orient='records'))140def format_user_prompt(df):141    return prompt_user_input.format_prompt(data=df.head().to_json(orient='records'))142 143def process_inputs(inputs) :144    return {'input_query': inputs['messages'].to_messages()[1]}145 146@traceable(process_inputs=process_inputs)147def run_llm(messages):148  try:149    response = llm.invoke(messages)150    print(response.content.replace("```", "'''").replace("json", ""))151    tests = json.loads(response.content.replace("```", "").replace("json", ""))152  except Exception as e:153      return e154  return tests155 156 157# Get Schema158def get_table_schema(table):159    result = conn.sql(f"SELECT sql, database_name, schema_name FROM duckdb_tables() where table_name ='{table}';").df()160    ddl_create = result.iloc[0,0]161    parent_database = result.iloc[0,1]162    schema_name = result.iloc[0,2]163    full_path = f"{parent_database}.{schema_name}.{table}"164    if schema_name != "main":165        old_path = f"{schema_name}.{table}"166    else:167        old_path = table168    ddl_create = ddl_create.replace(old_path, full_path)169    return full_path170 171def describe(df):172    173    numerical_info = pd.DataFrame()174    categorical_info = pd.DataFrame()175    if len(df.select_dtypes(include=['number']).columns) >= 1:176        numerical_info = df.select_dtypes(include=['number']).describe().T.reset_index()177        numerical_info.rename(columns={'index': 'column'}, inplace=True)178    if len(df.select_dtypes(include=['object']).columns) >= 1:179        categorical_info = df.select_dtypes(include=['object']).describe().T.reset_index()180        categorical_info.rename(columns={'index': 'column'}, inplace=True)181 182    return numerical_info, categorical_info183 184def validate_pandera(tests, df):185    validation_results = []186 187    for test in tests:188        column_name = test['column_name']189        try:190            rule = eval(test['pandera_rule'])  191            validated_column = rule(df[[column_name]])  192            validation_results.append({193            "Columns": column_name,194            "Result": "✅ Pass"195            })196        except Exception as e:197            validation_results.append({198            "Columns": column_name,199            "Result": f"❌ Fail - {str(e)}"200            })201    return pd.DataFrame(validation_results)202 203def statistics(df):204    profile = pp.ProfileReport(df)205    report_dict = profile.get_description()206    description, alerts = report_dict.table, report_dict.alerts207    # Statistics208    mapping = {209        'n': 'Number of observations',210        'n_var': 'Number of variables',211        'n_cells_missing': 'Number of cells missing',212        'n_vars_with_missing': 'Number of columns with missing data',213        'n_vars_all_missing': 'Columns with all missing data',214        'p_cells_missing': 'Missing cells (%)',215        'n_duplicates': 'Duplicated rows',216        'p_duplicates': 'Duplicated rows (%)',217    }218 219    updated_data = {mapping.get(k, k): v for k, v in description.items() if k != 'types'}220    # Add flattened types information221    if 'Text' in description.get('types', {}):222            updated_data['Number of text columns'] = description['types']['Text']223    if 'Categorical' in description.get('types', {}):224        updated_data['Number of categorical columns'] = description['types']['Categorical']225    if 'Numeric' in description.get('types', {}):226        updated_data['Number of numeric columns'] = description['types']['Numeric']227    if 'DateTime' in description.get('types', {}):228        updated_data['Number of datetime columns'] = description['types']['DateTime']229 230    df_statistics = pd.DataFrame(list(updated_data.items()), columns=['Statistic Description', 'Value'])231    df_statistics['Value'] = df_statistics['Value'].astype(int)232 233    # Alerts234    alerts_list = [(str(alert).replace('[', '').replace(']', ''), alert.alert_type_name) for alert in alerts]235    df_alerts = pd.DataFrame(alerts_list, columns=['Data Quality Issue', 'Category'])236 237    return df_statistics, df_alerts238#---------------------------------------239 240 241 242# Main Function243def main(table):244    schema = get_table_schema(table)245    246    # Create dlt pipeline247    table_name = create_pipeline(schema)248    249    # Load dlt pipeline250    connection, df = load_pipeline(table_name)251    252    # df = get_data_df(schema)253    df_statistics, df_alerts = statistics(df)254    describe_num, describe_cat  = describe(df)255   256    messages = format_prompt(df=df)257    tests = run_llm(messages)258   259    if isinstance(tests, Exception):260        tests = pd.DataFrame([{"error": f"❌ Unable to generate tests. {tests}"}])261        return df.head(10), df_statistics, df_alerts, describe_cat, describe_num, tests, pd.DataFrame([])262 263    tests_df = pd.DataFrame(tests)264    tests_df.rename(columns={tests_df.columns[0]: 'Column', tests_df.columns[1]: 'Rule Name', tests_df.columns[2]: 'Rules' }, inplace=True)265    pandera_results = validate_pandera(tests, df)266    267    connection.close()268    return df.head(10), df_statistics, df_alerts, describe_cat, describe_num, tests_df, pandera_results269 270def user_results(table, text_query):271    272    schema = get_table_schema(table)273    274    # Create dlt pipeline275    table_name = create_pipeline(schema)276    277    # Load dlt pipeline278    connection, df = load_pipeline(table_name)279    280    messages = format_user_prompt(df=df, user_description=text_query)281 282    print(f'Generated Tests from user input: {tests}')283    284    if isinstance(tests, Exception):285        tests = pd.DataFrame([{"error": f"❌ Unable to generate tests. {tests}"}])286        return tests, pd.DataFrame([])287 288    tests_df = pd.DataFrame(tests)289    tests_df.rename(columns={tests_df.columns[0]: 'Column', tests_df.columns[1]: 'Rule Name', tests_df.columns[2]: 'Rules' }, inplace=True)290    pandera_results = validate_pandera(tests, df)291    292    connection.close()293    294    return tests_df, pandera_results295    296# Custom CSS styling297custom_css = """298    print('Validated Tests with Pandera')299.gradio-container {300    background-color: #f0f4f8;301 302}303.logo {304    max-width: 200px;305    margin: 20px auto;306    display: block;307}308.gr-button {309    background-color: #4a90e2 !important;310}311.gr-button:hover {312    background-color: #3a7bc8 !important;313}314"""315 316with gr.Blocks(theme=gr.themes.Soft(primary_hue="purple", secondary_hue="indigo"), css=custom_css) as demo:317    gr.Image("logo.png", label=None, show_label=False, container=False, height=100)318 319    gr.Markdown("""320    <div style='text-align: center;'>321    <strong style='font-size: 36px;'>Dataset Test Workflow</strong>322    <br>323    <span style='font-size: 20px;'>Implement and Automate Data Validation Processes.</span>324    </div>325    """)326 327    with gr.Row():328        with gr.Column(scale=1):329            schema_dropdown = gr.Dropdown(choices=get_schemas(), label="Select Schema", interactive=True)330            tables_dropdown = gr.Dropdown(choices=[], label="Available Tables", value=None)331            with gr.Row():332                generate_result = gr.Button("Validate Data", variant="primary")333 334        with gr.Column(scale=2):335            with gr.Tabs():336 337                with gr.Tab("Description"):338                    with gr.Row():339                        with gr.Column():340                            data_description = gr.DataFrame(label="Data Description", value=[], interactive=False)341                    with gr.Row():342                        with gr.Column():343                            describe_cat = gr.DataFrame(label="Categorical Information", value=[], interactive=False)344                        with gr.Column():    345                            describe_num = gr.DataFrame(label="Numerical Information", value=[], interactive=False)346 347                with gr.Tab("Alerts"):348                    data_alerts = gr.DataFrame(label="Alerts", value=[], interactive=False)349 350                with gr.Tab("Rules & Validations"):351                    tests_output = gr.DataFrame(label="Validation Rules", value=[], interactive=False)352                    test_result_output = gr.DataFrame(label="Validation Result", value=[], interactive=False)353                354                with gr.Tab("Data"):355                    result_output = gr.DataFrame(label="Dataframe (10 Rows)", value=[], interactive=False)356                357                with gr.Tab('Text to Validation'):358                    with gr.Row():359                        query_input = gr.Textbox(lines=5, label="Text Query", placeholder="Enter Text Query to Generate Validation e.g. Validate that the incident_zip column contains valid 5-digit ZIP codes.")360                    with gr.Row():361                        with gr.Column():  362                            pass363                        with gr.Column(scale=1, min_width=50):  364                            user_generate_result = gr.Button("Validate Data", variant="primary" )  365                366                    with gr.Row():367                        with gr.Column():368                            query_tests = gr.DataFrame(label="Validation Rules", value=[], interactive=False)369                        with gr.Column():370                            query_result = gr.DataFrame(label="Validation Result", value=[], interactive=False)371                                 372        schema_dropdown.change(update_table_names, inputs=schema_dropdown, outputs=tables_dropdown)373        generate_result.click(main, inputs=[tables_dropdown], outputs=[result_output, data_description, data_alerts, describe_cat, describe_num, tests_output, test_result_output])374        user_generate_result.click(user_results, inputs=[tables_dropdown, query_input], outputs=[query_tests, query_result])375 376if __name__ == "__main__":377    demo.launch(debug=True)378 379 380