{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os, sys\nfrom graphnet.data.sqlite.sqlite_utilities import create_table\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\nimport sqlite3\nimport pyarrow.parquet as pq\nimport sqlalchemy\nfrom tqdm import tqdm\nfrom typing import Any, Dict, List, Optional\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport pickle\nimport time\nimport gc\n\ninput_data_folder = \"./data/train\"\nmeta_data_path = \"./data/train_meta.parquet\"\ngeometry_table = pd.read_csv(\"./data/sensor_geometry.csv\")\n\n\ndef load_input(batch_id: int, input_data_folder: str, event_ids: []) -> pd.DataFrame:\n    detector_readings = pd.read_parquet(\n        path=f\"{input_data_folder}/batch_{batch_id}.parquet\"\n    )\n    detector_readings = detector_readings.loc[detector_readings.index.isin(event_ids)]\n    sensor_positions = geometry_table.loc[\n        detector_readings[\"sensor_id\"], [\"x\", \"y\", \"z\"]\n    ]\n    sensor_positions.index = detector_readings.index\n    detector_readings_copy = detector_readings.copy()\n    detector_readings_copy.loc[:, \"x\"] = sensor_positions[\"x\"]\n    detector_readings_copy.loc[:, \"y\"] = sensor_positions[\"y\"]\n    detector_readings_copy.loc[:, \"z\"] = sensor_positions[\"z\"]\n    detector_readings = detector_readings_copy\n    del detector_readings_copy\n    detector_readings[\"auxiliary\"].replace({True: 1, False: 0}, inplace=True)\n\n    return detector_readings.reset_index()\n\n\ndef add_to_table(\n    database_path: str,\n    df: pd.DataFrame,\n    table_name: str,\n    is_primary_key: bool,\n    engine: sqlalchemy.engine.base.Engine,\n) -> None:\n    try:\n        create_table(\n            columns=df.columns,\n            database_path=database_path,\n            table_name=table_name,\n            integer_primary_key=is_primary_key,\n            index_column=\"event_id\",\n        )\n    except sqlite3.OperationalError as e:\n        if \"already exists\" in str(e):\n            pass\n        else:\n            raise e\n\n    df.to_sql(table_name, con=engine, index=False, if_exists=\"append\", chunksize=200000)\n    engine.dispose()\n    return\n\n\ndef convert_to_sqlite(\n    meta_data_path: str,\n    database_path: str,\n    input_data_folder: str,\n    batch_size: int = 200000,\n    batch_ids: list = [],\n    event_ids: list = [],\n    engine: sqlalchemy.engine.base.Engine = None,\n) -> None:\n    max_pulse = 600\n    meta_data_iter = pq.ParquetFile(meta_data_path).iter_batches(batch_size=batch_size)\n    batch_id = 1\n    for meta_data_batch_raw in tqdm(meta_data_iter):\n        if batch_id in batch_ids:\n            meta_data_batch = meta_data_batch_raw.to_pandas()\n            meta_data_batch.drop(\n                columns=[\"first_pulse_index\", \"last_pulse_index\"], inplace=True\n            )\n            meta_data_batch = meta_data_batch.loc[\n                meta_data_batch[\"event_id\"].isin(event_ids)\n            ].reset_index(drop=True)\n\n            if meta_data_batch.shape[0] > 0:\n                pulses = load_input(\n                    batch_id=batch_id,\n                    input_data_folder=input_data_folder,\n                    event_ids=event_ids,\n                )\n                pulses = (\n                    pulses.groupby(\"event_id\").head(max_pulse).reset_index(drop=True)\n                )\n\n                first_event_id_meta = meta_data_batch.iloc[0].event_id\n                last_event_id_meta = meta_data_batch.iloc[-1].event_id\n                if batch_id == 1:\n                    add_to_table(\n                        database_path=database_path,\n                        df=meta_data_batch,\n                        table_name=\"meta_table\",\n                        is_primary_key=True,\n                        engine=engine,\n                    )\n                else:\n                    with sqlite3.connect(database_path) as con:\n                        query = f\"select event_id from meta_table where event_id in ({first_event_id_meta}, {last_event_id_meta})\"\n                        events_meta_df = pd.read_sql(query, con)\n                    if events_meta_df.shape[0] == 0:\n                        add_to_table(\n                            database_path=database_path,\n                            df=meta_data_batch,\n                            table_name=\"meta_table\",\n                            is_primary_key=True,\n                            engine=engine,\n                        )\n\n                first_event_id_pulses = pulses.iloc[0].event_id\n                last_event_id_pulses = pulses.iloc[-1].event_id\n                if batch_id == 1:\n                    add_to_table(\n                        database_path=database_path,\n                        df=pulses,\n                        table_name=\"pulse_table\",\n                        is_primary_key=False,\n                        engine=engine,\n                    )\n                else:\n                    with sqlite3.connect(database_path) as con:\n                        query = f\"select event_id from pulse_table where event_id in ({first_event_id_pulses}, {last_event_id_pulses})\"\n                        events_pulse_df = pd.read_sql(query, con)\n                    if events_pulse_df.shape[0] == 0:\n                        add_to_table(\n                            database_path=database_path,\n                            df=pulses,\n                            table_name=\"pulse_table\",\n                            is_primary_key=False,\n                            engine=engine,\n                        )\n                        gc.collect()\n                batch_id += 1\n    del meta_data_iter\n\n\nwith open(\"focus_dict.pkl\", \"rb\") as f:\n    focus_dict = pickle.load(f)\n\nidx = 1\nevent_id_list = focus_dict[f\"f{idx}\"]\ndatabase_path = f\"./data/F{idx}/focus_batch_{idx}.db\"\nengine = sqlalchemy.create_engine(\"sqlite:///\" + database_path)\nconvert_to_sqlite(\n    meta_data_path,\n    database_path=database_path,\n    input_data_folder=input_data_folder,\n    batch_size=200000,\n    batch_ids=list(range(1, 661, 1)),\n    event_ids=event_id_list,\n    engine=engine,\n)\n\nwith sqlite3.connect(database_path) as con:\n    query = \"select event_id from meta_table\"\n    events_df = pd.read_sql(query, con)\n\ntrain_selection, validate_selection = train_test_split(\n    np.arange(0, events_df.shape[0], 1), shuffle=True, random_state=42, test_size=0.02\n)\n\ntrain_selection_events = events_df[events_df.index.isin(train_selection)][\n    \"event_id\"\n].to_list()\nvalidate_selection_events = events_df[events_df.index.isin(validate_selection)][\n    \"event_id\"\n].to_list()\nevent_dict = {\"train\": train_selection_events, \"validate\": validate_selection_events}\nwith open(f\"data/F{idx}/event_dict.pkl\", \"wb\") as f:\n    pickle.dump(event_dict, f)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}