cehrbert-data 0.0.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- __init__.py +0 -0
- cehrbert_data/__init__.py +0 -0
- cehrbert_data/apps/__init__.py +0 -0
- cehrbert_data/apps/generate_concept_similarity_table.py +423 -0
- cehrbert_data/apps/generate_hierarchical_bert_training_data.py +238 -0
- cehrbert_data/apps/generate_included_concept_list.py +116 -0
- cehrbert_data/apps/generate_information_content.py +131 -0
- cehrbert_data/apps/generate_required_labs.py +112 -0
- cehrbert_data/apps/generate_training_data.py +337 -0
- cehrbert_data/cohorts/__init__.py +0 -0
- cehrbert_data/cohorts/atrial_fibrillation.py +45 -0
- cehrbert_data/cohorts/cabg.py +72 -0
- cehrbert_data/cohorts/coronary_artery_disease.py +84 -0
- cehrbert_data/cohorts/covid.py +43 -0
- cehrbert_data/cohorts/covid_inpatient.py +84 -0
- cehrbert_data/cohorts/death.py +46 -0
- cehrbert_data/cohorts/heart_failure.py +423 -0
- cehrbert_data/cohorts/ischemic_stroke.py +45 -0
- cehrbert_data/cohorts/last_visit_discharged_home.py +35 -0
- cehrbert_data/cohorts/query_builder.py +153 -0
- cehrbert_data/cohorts/spark_app_base.py +745 -0
- cehrbert_data/cohorts/type_two_diabietes.py +166 -0
- cehrbert_data/cohorts/ventilation.py +22 -0
- cehrbert_data/config/__init__.py +0 -0
- cehrbert_data/config/output_names.py +9 -0
- cehrbert_data/const/__init__.py +0 -0
- cehrbert_data/const/__pycache__/__init__.cpython-311.pyc +0 -0
- cehrbert_data/const/__pycache__/common.cpython-311.pyc +0 -0
- cehrbert_data/const/common.py +28 -0
- cehrbert_data/decorators/__init__.py +0 -0
- cehrbert_data/decorators/__pycache__/__init__.cpython-311.pyc +0 -0
- cehrbert_data/decorators/__pycache__/patient_event_decorator.cpython-311.pyc +0 -0
- cehrbert_data/decorators/patient_event_decorator.py +759 -0
- cehrbert_data/prediction_cohorts/__init__.py +0 -0
- cehrbert_data/prediction_cohorts/afib_ischemic_stroke.py +14 -0
- cehrbert_data/prediction_cohorts/cad_cabg_cohort.py +19 -0
- cehrbert_data/prediction_cohorts/cad_hf_cohort.py +14 -0
- cehrbert_data/prediction_cohorts/copd_readmission.py +76 -0
- cehrbert_data/prediction_cohorts/covid_death.py +14 -0
- cehrbert_data/prediction_cohorts/covid_ventilation.py +14 -0
- cehrbert_data/prediction_cohorts/discharge_home_death.py +18 -0
- cehrbert_data/prediction_cohorts/hf_readmission.py +82 -0
- cehrbert_data/prediction_cohorts/hospitalization.py +100 -0
- cehrbert_data/prediction_cohorts/hospitalization_mortality.py +77 -0
- cehrbert_data/prediction_cohorts/t2dm_hf_cohort.py +14 -0
- cehrbert_data/queries/__init__.py +0 -0
- cehrbert_data/queries/measurement_unit_stats_query.py +42 -0
- cehrbert_data/tools/__init__.py +0 -0
- cehrbert_data/tools/download_omop_tables.py +141 -0
- cehrbert_data/utils/__init__.py +0 -0
- cehrbert_data/utils/spark_parse_args.py +350 -0
- cehrbert_data/utils/spark_utils.py +1405 -0
- cehrbert_data-0.0.1.dist-info/LICENSE +21 -0
- cehrbert_data-0.0.1.dist-info/METADATA +116 -0
- cehrbert_data-0.0.1.dist-info/RECORD +57 -0
- cehrbert_data-0.0.1.dist-info/WHEEL +5 -0
- cehrbert_data-0.0.1.dist-info/top_level.txt +2 -0
|
@@ -0,0 +1,116 @@
|
|
|
1
|
+
"""
|
|
2
|
+
This module generates a qualified concept list by processing patient event data across various.
|
|
3
|
+
|
|
4
|
+
domain tables (e.g., condition_occurrence, procedure_occurrence, drug_exposure) and applying a
|
|
5
|
+
patient frequency filter to retain concepts linked to a minimum number of patients.
|
|
6
|
+
|
|
7
|
+
Key Functions:
|
|
8
|
+
- preprocess_domain_table: Preprocesses domain tables to prepare for event extraction.
|
|
9
|
+
- join_domain_tables: Joins multiple domain tables into a unified DataFrame.
|
|
10
|
+
- main: Coordinates the entire process of reading domain tables, applying frequency filters,
|
|
11
|
+
and saving the qualified concept list.
|
|
12
|
+
|
|
13
|
+
Command-line Arguments:
|
|
14
|
+
- input_folder: Directory containing the input data.
|
|
15
|
+
- output_folder: Directory where the qualified concept list will be saved.
|
|
16
|
+
- min_num_of_patients: Minimum number of patients linked to a concept for it to be included.
|
|
17
|
+
- with_drug_rollup: Boolean flag indicating whether drug concept rollups should be applied.
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
import os
|
|
21
|
+
|
|
22
|
+
from pyspark.sql import SparkSession
|
|
23
|
+
from pyspark.sql import functions as F
|
|
24
|
+
|
|
25
|
+
from cehrbert_data.config.output_names import QUALIFIED_CONCEPT_LIST_PATH
|
|
26
|
+
from cehrbert_data.const.common import MEASUREMENT
|
|
27
|
+
from cehrbert_data.utils.spark_utils import join_domain_tables, preprocess_domain_table
|
|
28
|
+
|
|
29
|
+
DOMAIN_TABLE_LIST = ["condition_occurrence", "procedure_occurrence", "drug_exposure"]
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def main(input_folder, output_folder, min_num_of_patients, with_drug_rollup: bool = True):
|
|
33
|
+
"""
|
|
34
|
+
Main function to generate a qualified concept list based on patient event data from multiple.
|
|
35
|
+
|
|
36
|
+
domain tables.
|
|
37
|
+
|
|
38
|
+
Args:
|
|
39
|
+
input_folder (str): The directory where the input data is stored.
|
|
40
|
+
output_folder (str): The directory where the output (qualified concept list) will be saved.
|
|
41
|
+
min_num_of_patients (int): Minimum number of patients that a concept must be linked to for
|
|
42
|
+
nclusion.
|
|
43
|
+
with_drug_rollup (bool): If True, applies drug rollup logic to the drug_exposure domain.
|
|
44
|
+
|
|
45
|
+
The function processes patient event data across various domain tables, excludes low-frequency
|
|
46
|
+
concepts, and saves the filtered concepts to a specified output folder.
|
|
47
|
+
"""
|
|
48
|
+
spark = SparkSession.builder.appName("Generate concept list").getOrCreate()
|
|
49
|
+
|
|
50
|
+
domain_tables = []
|
|
51
|
+
# Exclude measurement from domain_table_list if exists because we need to process measurement
|
|
52
|
+
# in a different way
|
|
53
|
+
for domain_table_name in DOMAIN_TABLE_LIST:
|
|
54
|
+
if domain_table_name != MEASUREMENT:
|
|
55
|
+
domain_tables.append(
|
|
56
|
+
preprocess_domain_table(
|
|
57
|
+
spark,
|
|
58
|
+
input_folder,
|
|
59
|
+
domain_table_name,
|
|
60
|
+
with_drug_rollup=with_drug_rollup,
|
|
61
|
+
)
|
|
62
|
+
)
|
|
63
|
+
|
|
64
|
+
# Union all domain table records
|
|
65
|
+
patient_events = join_domain_tables(domain_tables)
|
|
66
|
+
|
|
67
|
+
# Filter out concepts that are linked to less than 100 patients
|
|
68
|
+
qualified_concepts = (
|
|
69
|
+
patient_events.where("visit_occurrence_id IS NOT NULL")
|
|
70
|
+
.groupBy("standard_concept_id")
|
|
71
|
+
.agg(F.countDistinct("person_id").alias("freq"))
|
|
72
|
+
.where(F.col("freq") >= min_num_of_patients)
|
|
73
|
+
)
|
|
74
|
+
|
|
75
|
+
qualified_concepts.write.mode("overwrite").parquet(os.path.join(output_folder, QUALIFIED_CONCEPT_LIST_PATH))
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
if __name__ == "__main__":
|
|
79
|
+
import argparse
|
|
80
|
+
|
|
81
|
+
parser = argparse.ArgumentParser(description="Arguments for generate concept list to be included")
|
|
82
|
+
parser.add_argument(
|
|
83
|
+
"-i",
|
|
84
|
+
"--input_folder",
|
|
85
|
+
dest="input_folder",
|
|
86
|
+
action="store",
|
|
87
|
+
help="The path for your input_folder where the raw data is",
|
|
88
|
+
required=True,
|
|
89
|
+
)
|
|
90
|
+
parser.add_argument(
|
|
91
|
+
"-o",
|
|
92
|
+
"--output_folder",
|
|
93
|
+
dest="output_folder",
|
|
94
|
+
action="store",
|
|
95
|
+
help="The path for your output_folder",
|
|
96
|
+
required=True,
|
|
97
|
+
)
|
|
98
|
+
parser.add_argument(
|
|
99
|
+
"--min_num_of_patients",
|
|
100
|
+
dest="min_num_of_patients",
|
|
101
|
+
action="store",
|
|
102
|
+
type=int,
|
|
103
|
+
default=0,
|
|
104
|
+
help="Min no.of patients linked to concepts to be included",
|
|
105
|
+
required=False,
|
|
106
|
+
)
|
|
107
|
+
parser.add_argument("--with_drug_rollup", dest="with_drug_rollup", action="store_true")
|
|
108
|
+
|
|
109
|
+
ARGS = parser.parse_args()
|
|
110
|
+
|
|
111
|
+
main(
|
|
112
|
+
ARGS.input_folder,
|
|
113
|
+
ARGS.output_folder,
|
|
114
|
+
ARGS.min_num_of_patients,
|
|
115
|
+
ARGS.with_drug_rollup,
|
|
116
|
+
)
|
|
@@ -0,0 +1,131 @@
|
|
|
1
|
+
"""
|
|
2
|
+
This module generates an information content table based on a list of domain tables from OMOP data.
|
|
3
|
+
|
|
4
|
+
It processes patient event data, calculates the frequency of each concept, and computes information
|
|
5
|
+
conten using the concept ancestor hierarchy. The results are written to a specified output path.
|
|
6
|
+
|
|
7
|
+
Key Functions:
|
|
8
|
+
- preprocess_domain_table: Preprocess the domain tables for analysis.
|
|
9
|
+
- join_domain_tables: Join multiple domain tables to generate a unified patient event table.
|
|
10
|
+
- main: Orchestrates the process of reading input data, calculating concept frequencies,
|
|
11
|
+
and generating the information content table.
|
|
12
|
+
|
|
13
|
+
Command-line Arguments:
|
|
14
|
+
- input_folder: The folder containing the raw OMOP domain data.
|
|
15
|
+
- output_folder: The folder where the results will be stored.
|
|
16
|
+
- domain_table_list: A list of OMOP domain tables to include in the analysis.
|
|
17
|
+
- date_filter: Optional date filter for processing the data.
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
import datetime
|
|
21
|
+
import logging
|
|
22
|
+
import os
|
|
23
|
+
|
|
24
|
+
from pyspark.sql import SparkSession
|
|
25
|
+
from pyspark.sql import functions as F
|
|
26
|
+
|
|
27
|
+
from cehrbert_data.config.output_names import INFORMATION_CONTENT_DATA_PATH
|
|
28
|
+
from cehrbert_data.const.common import CONCEPT_ANCESTOR
|
|
29
|
+
from cehrbert_data.utils.spark_utils import join_domain_tables, preprocess_domain_table, validate_table_names
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def main(input_folder, output_folder, domain_table_list, date_filter):
|
|
33
|
+
"""Create the information content table.
|
|
34
|
+
|
|
35
|
+
Keyword arguments:
|
|
36
|
+
domain_tables -- the array containing the OMOP domain tables except visit_occurrence
|
|
37
|
+
concept_id_frequency_output -- the path for writing the concept frequency output
|
|
38
|
+
|
|
39
|
+
This function creates the information content table based on the given domain tables
|
|
40
|
+
"""
|
|
41
|
+
|
|
42
|
+
spark = SparkSession.builder.appName("Generate the information content table").getOrCreate()
|
|
43
|
+
|
|
44
|
+
logger = logging.getLogger(__name__)
|
|
45
|
+
logger.info(
|
|
46
|
+
"input_folder: %s\noutput_folder: %s\ndomain_table_list: %s\ndate_filter: %s",
|
|
47
|
+
input_folder,
|
|
48
|
+
output_folder,
|
|
49
|
+
domain_table_list,
|
|
50
|
+
date_filter,
|
|
51
|
+
)
|
|
52
|
+
|
|
53
|
+
concept_ancestor = preprocess_domain_table(spark, input_folder, CONCEPT_ANCESTOR)
|
|
54
|
+
domain_tables = []
|
|
55
|
+
for domain_table_name in domain_table_list:
|
|
56
|
+
domain_tables.append(preprocess_domain_table(spark, input_folder, domain_table_name))
|
|
57
|
+
|
|
58
|
+
patient_events = join_domain_tables(domain_tables)
|
|
59
|
+
|
|
60
|
+
# Remove all concept_id records
|
|
61
|
+
patient_events = patient_events.where("standard_concept_id <> 0")
|
|
62
|
+
|
|
63
|
+
# Get the total count
|
|
64
|
+
total_count = patient_events.distinct().count()
|
|
65
|
+
|
|
66
|
+
# Count the frequency of each concept
|
|
67
|
+
concept_frequency = patient_events.distinct().groupBy("standard_concept_id").count()
|
|
68
|
+
|
|
69
|
+
# left join b/w descendent_concept_id and the standard_concept_id in the concept freq table
|
|
70
|
+
freq_df = (
|
|
71
|
+
concept_frequency.join(
|
|
72
|
+
concept_ancestor,
|
|
73
|
+
F.col("descendant_concept_id") == F.col("standard_concept_id"),
|
|
74
|
+
)
|
|
75
|
+
.groupBy("ancestor_concept_id")
|
|
76
|
+
.sum("count")
|
|
77
|
+
.withColumnRenamed("ancestor_concept_id", "concept_id")
|
|
78
|
+
.withColumnRenamed("sum(count)", "count")
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
# Calculate information content for each concept
|
|
82
|
+
information_content = freq_df.withColumn("information_content", (-F.log(F.col("count") / total_count))).withColumn(
|
|
83
|
+
"probability", F.col("count") / total_count
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
information_content.write.mode("overwrite").parquet(os.path.join(output_folder, INFORMATION_CONTENT_DATA_PATH))
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
if __name__ == "__main__":
|
|
90
|
+
import argparse
|
|
91
|
+
|
|
92
|
+
parser = argparse.ArgumentParser(description="Arguments for generate training data for Bert")
|
|
93
|
+
parser.add_argument(
|
|
94
|
+
"-i",
|
|
95
|
+
"--input_folder",
|
|
96
|
+
dest="input_folder",
|
|
97
|
+
action="store",
|
|
98
|
+
help="The path for your input_folder where the raw data is",
|
|
99
|
+
required=True,
|
|
100
|
+
)
|
|
101
|
+
parser.add_argument(
|
|
102
|
+
"-o",
|
|
103
|
+
"--output_folder",
|
|
104
|
+
dest="output_folder",
|
|
105
|
+
action="store",
|
|
106
|
+
help="The path for your output_folder",
|
|
107
|
+
required=True,
|
|
108
|
+
)
|
|
109
|
+
parser.add_argument(
|
|
110
|
+
"-tc",
|
|
111
|
+
"--domain_table_list",
|
|
112
|
+
dest="domain_table_list",
|
|
113
|
+
nargs="+",
|
|
114
|
+
action="store",
|
|
115
|
+
help="The list of domain tables you want to download",
|
|
116
|
+
type=validate_table_names,
|
|
117
|
+
required=True,
|
|
118
|
+
)
|
|
119
|
+
parser.add_argument(
|
|
120
|
+
"-d",
|
|
121
|
+
"--date_filter",
|
|
122
|
+
dest="date_filter",
|
|
123
|
+
type=lambda s: datetime.datetime.strptime(s, "%Y-%m-%d"),
|
|
124
|
+
action="store",
|
|
125
|
+
required=False,
|
|
126
|
+
default="2018-01-01",
|
|
127
|
+
)
|
|
128
|
+
|
|
129
|
+
ARGS = parser.parse_args()
|
|
130
|
+
|
|
131
|
+
main(ARGS.input_folder, ARGS.output_folder, ARGS.domain_table_list, ARGS.date_filter)
|
|
@@ -0,0 +1,112 @@
|
|
|
1
|
+
import argparse
|
|
2
|
+
import os
|
|
3
|
+
|
|
4
|
+
from pyspark.sql import SparkSession
|
|
5
|
+
from pyspark.sql import functions as F
|
|
6
|
+
from pyspark.sql.window import Window
|
|
7
|
+
|
|
8
|
+
from cehrbert_data.const.common import CONCEPT, MEASUREMENT, REQUIRED_MEASUREMENT
|
|
9
|
+
from cehrbert_data.utils.spark_utils import preprocess_domain_table
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def main(input_folder, output_folder, num_of_numeric_labs, num_of_categorical_labs):
|
|
13
|
+
spark = SparkSession.builder.appName("Generate required labs").getOrCreate()
|
|
14
|
+
|
|
15
|
+
# Load measurement as a dataframe in pyspark
|
|
16
|
+
measurement = preprocess_domain_table(spark, input_folder, MEASUREMENT)
|
|
17
|
+
concept = preprocess_domain_table(spark, input_folder, CONCEPT)
|
|
18
|
+
|
|
19
|
+
# Create the local measurement view
|
|
20
|
+
measurement.createOrReplaceTempView("measurement")
|
|
21
|
+
|
|
22
|
+
# Create the local concept view
|
|
23
|
+
concept.createOrReplaceTempView("concept")
|
|
24
|
+
|
|
25
|
+
popular_labs = spark.sql(
|
|
26
|
+
"""
|
|
27
|
+
SELECT
|
|
28
|
+
m.measurement_concept_id,
|
|
29
|
+
c.concept_name,
|
|
30
|
+
COUNT(*) AS freq,
|
|
31
|
+
SUM(CASE WHEN m.value_as_number IS NOT NULL THEN 1 ELSE 0 END) / COUNT(*) AS numeric_percentage,
|
|
32
|
+
SUM(CASE WHEN m.value_as_concept_id IS NOT NULL AND m.value_as_concept_id <> 0 THEN 1 ELSE 0 END) / COUNT(*) AS categorical_percentage
|
|
33
|
+
FROM measurement AS m
|
|
34
|
+
JOIN concept AS c
|
|
35
|
+
ON m.measurement_concept_id = c.concept_id
|
|
36
|
+
WHERE m.measurement_concept_id <> 0
|
|
37
|
+
GROUP BY m.measurement_concept_id, c.concept_name
|
|
38
|
+
ORDER BY COUNT(*) DESC
|
|
39
|
+
"""
|
|
40
|
+
)
|
|
41
|
+
|
|
42
|
+
# Cache the dataframe for faster computation in the below transformations
|
|
43
|
+
popular_labs.cache()
|
|
44
|
+
|
|
45
|
+
popular_numeric_labs = (
|
|
46
|
+
popular_labs.withColumn("is_numeric", F.col("numeric_percentage") >= 0.5)
|
|
47
|
+
.where("is_numeric")
|
|
48
|
+
.withColumn("rn", F.row_number().over(Window.orderBy(F.desc("freq"))))
|
|
49
|
+
.where(F.col("rn") <= num_of_numeric_labs)
|
|
50
|
+
.drop("rn")
|
|
51
|
+
)
|
|
52
|
+
|
|
53
|
+
popular_categorical_labs = (
|
|
54
|
+
popular_labs.withColumn("is_categorical", F.col("categorical_percentage") >= 0.5)
|
|
55
|
+
.where("is_categorical")
|
|
56
|
+
.withColumn("is_numeric", ~F.col("is_categorical"))
|
|
57
|
+
.withColumn("rn", F.row_number().over(Window.orderBy(F.desc("freq"))))
|
|
58
|
+
.where(F.col("rn") <= num_of_categorical_labs)
|
|
59
|
+
.drop("is_categorical")
|
|
60
|
+
.drop("rn")
|
|
61
|
+
)
|
|
62
|
+
|
|
63
|
+
popular_numeric_labs.unionAll(popular_categorical_labs).write.mode("overwrite").parquet(
|
|
64
|
+
os.path.join(output_folder, REQUIRED_MEASUREMENT)
|
|
65
|
+
)
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
if __name__ == "__main__":
|
|
69
|
+
parser = argparse.ArgumentParser(description="Arguments for generate " "required labs to be included")
|
|
70
|
+
parser.add_argument(
|
|
71
|
+
"-i",
|
|
72
|
+
"--input_folder",
|
|
73
|
+
dest="input_folder",
|
|
74
|
+
action="store",
|
|
75
|
+
help="The path for your input_folder where the raw data is",
|
|
76
|
+
required=True,
|
|
77
|
+
)
|
|
78
|
+
parser.add_argument(
|
|
79
|
+
"-o",
|
|
80
|
+
"--output_folder",
|
|
81
|
+
dest="output_folder",
|
|
82
|
+
action="store",
|
|
83
|
+
help="The path for your output_folder",
|
|
84
|
+
required=True,
|
|
85
|
+
)
|
|
86
|
+
parser.add_argument(
|
|
87
|
+
"--num_of_numeric_labs",
|
|
88
|
+
dest="num_of_numeric_labs",
|
|
89
|
+
action="store",
|
|
90
|
+
type=int,
|
|
91
|
+
default=100,
|
|
92
|
+
help="The top most popular numeric labs to be included",
|
|
93
|
+
required=False,
|
|
94
|
+
)
|
|
95
|
+
parser.add_argument(
|
|
96
|
+
"--num_of_categorical_labs",
|
|
97
|
+
dest="num_of_categorical_labs",
|
|
98
|
+
action="store",
|
|
99
|
+
type=int,
|
|
100
|
+
default=100,
|
|
101
|
+
help="The top most popular categorical labs to be included",
|
|
102
|
+
required=False,
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
ARGS = parser.parse_args()
|
|
106
|
+
|
|
107
|
+
main(
|
|
108
|
+
ARGS.input_folder,
|
|
109
|
+
ARGS.output_folder,
|
|
110
|
+
ARGS.num_of_numeric_labs,
|
|
111
|
+
ARGS.num_of_categorical_labs,
|
|
112
|
+
)
|
|
@@ -0,0 +1,337 @@
|
|
|
1
|
+
import argparse
|
|
2
|
+
import datetime
|
|
3
|
+
import logging
|
|
4
|
+
import os
|
|
5
|
+
import shutil
|
|
6
|
+
|
|
7
|
+
from pyspark.sql import SparkSession
|
|
8
|
+
from pyspark.sql import functions as F
|
|
9
|
+
from pyspark.sql.window import Window
|
|
10
|
+
|
|
11
|
+
from cehrbert_data.const.common import DEATH, MEASUREMENT, PERSON, REQUIRED_MEASUREMENT, VISIT_OCCURRENCE
|
|
12
|
+
from cehrbert_data.decorators.patient_event_decorator import AttType
|
|
13
|
+
from cehrbert_data.utils.spark_utils import (
|
|
14
|
+
create_sequence_data,
|
|
15
|
+
create_sequence_data_with_att,
|
|
16
|
+
join_domain_tables,
|
|
17
|
+
preprocess_domain_table,
|
|
18
|
+
process_measurement,
|
|
19
|
+
validate_table_names,
|
|
20
|
+
)
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def main(
|
|
24
|
+
input_folder,
|
|
25
|
+
output_folder,
|
|
26
|
+
domain_table_list,
|
|
27
|
+
date_filter,
|
|
28
|
+
include_visit_type,
|
|
29
|
+
is_new_patient_representation,
|
|
30
|
+
exclude_visit_tokens,
|
|
31
|
+
is_classic_bert,
|
|
32
|
+
include_prolonged_stay,
|
|
33
|
+
include_concept_list: bool,
|
|
34
|
+
gpt_patient_sequence: bool,
|
|
35
|
+
apply_age_filter: bool,
|
|
36
|
+
include_death: bool,
|
|
37
|
+
att_type: AttType,
|
|
38
|
+
include_sequence_information_content: bool = False,
|
|
39
|
+
exclude_demographic: bool = False,
|
|
40
|
+
use_age_group: bool = False,
|
|
41
|
+
with_drug_rollup: bool = True,
|
|
42
|
+
include_inpatient_hour_token: bool = False,
|
|
43
|
+
continue_from_events: bool = False,
|
|
44
|
+
):
|
|
45
|
+
spark = SparkSession.builder.appName("Generate CEHR-BERT Training Data").getOrCreate()
|
|
46
|
+
|
|
47
|
+
logger = logging.getLogger(__name__)
|
|
48
|
+
logger.info(
|
|
49
|
+
f"input_folder: {input_folder}\n"
|
|
50
|
+
f"output_folder: {output_folder}\n"
|
|
51
|
+
f"domain_table_list: {domain_table_list}\n"
|
|
52
|
+
f"date_filter: {date_filter}\n"
|
|
53
|
+
f"include_visit_type: {include_visit_type}\n"
|
|
54
|
+
f"is_new_patient_representation: {is_new_patient_representation}\n"
|
|
55
|
+
f"exclude_visit_tokens: {exclude_visit_tokens}\n"
|
|
56
|
+
f"is_classic_bert: {is_classic_bert}\n"
|
|
57
|
+
f"include_prolonged_stay: {include_prolonged_stay}\n"
|
|
58
|
+
f"include_concept_list: {include_concept_list}\n"
|
|
59
|
+
f"gpt_patient_sequence: {gpt_patient_sequence}\n"
|
|
60
|
+
f"apply_age_filter: {apply_age_filter}\n"
|
|
61
|
+
f"include_death: {include_death}\n"
|
|
62
|
+
f"att_type: {att_type}\n"
|
|
63
|
+
f"exclude_demographic: {exclude_demographic}\n"
|
|
64
|
+
f"use_age_group: {use_age_group}\n"
|
|
65
|
+
f"with_drug_rollup: {with_drug_rollup}\n"
|
|
66
|
+
)
|
|
67
|
+
|
|
68
|
+
domain_tables = []
|
|
69
|
+
for domain_table_name in domain_table_list:
|
|
70
|
+
if domain_table_name != MEASUREMENT:
|
|
71
|
+
domain_tables.append(
|
|
72
|
+
preprocess_domain_table(
|
|
73
|
+
spark,
|
|
74
|
+
input_folder,
|
|
75
|
+
domain_table_name,
|
|
76
|
+
with_drug_rollup=with_drug_rollup,
|
|
77
|
+
)
|
|
78
|
+
)
|
|
79
|
+
|
|
80
|
+
visit_occurrence = preprocess_domain_table(spark, input_folder, VISIT_OCCURRENCE)
|
|
81
|
+
visit_occurrence = visit_occurrence.select(
|
|
82
|
+
"visit_occurrence_id",
|
|
83
|
+
"visit_start_date",
|
|
84
|
+
"visit_start_datetime",
|
|
85
|
+
"visit_end_date",
|
|
86
|
+
"visit_concept_id",
|
|
87
|
+
"person_id",
|
|
88
|
+
"discharged_to_concept_id",
|
|
89
|
+
)
|
|
90
|
+
person = preprocess_domain_table(spark, input_folder, PERSON)
|
|
91
|
+
birth_datetime_udf = F.coalesce("birth_datetime", F.concat("year_of_birth", F.lit("-01-01")).cast("timestamp"))
|
|
92
|
+
person = person.select(
|
|
93
|
+
"person_id",
|
|
94
|
+
birth_datetime_udf.alias("birth_datetime"),
|
|
95
|
+
"race_concept_id",
|
|
96
|
+
"gender_concept_id",
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
visit_occurrence_person = visit_occurrence.join(person, "person_id").withColumn(
|
|
100
|
+
"age",
|
|
101
|
+
F.ceil(F.months_between(F.col("visit_start_date"), F.col("birth_datetime")) / F.lit(12)),
|
|
102
|
+
)
|
|
103
|
+
visit_occurrence_person = visit_occurrence_person.drop("birth_datetime")
|
|
104
|
+
|
|
105
|
+
death = preprocess_domain_table(spark, input_folder, DEATH) if include_death else None
|
|
106
|
+
|
|
107
|
+
patient_events = join_domain_tables(domain_tables)
|
|
108
|
+
|
|
109
|
+
if include_concept_list and patient_events:
|
|
110
|
+
column_names = patient_events.schema.fieldNames()
|
|
111
|
+
# Filter out concepts
|
|
112
|
+
qualified_concepts = preprocess_domain_table(spark, input_folder, "qualified_concept_list").select(
|
|
113
|
+
"standard_concept_id"
|
|
114
|
+
)
|
|
115
|
+
|
|
116
|
+
patient_events = patient_events.join(qualified_concepts, "standard_concept_id").select(column_names)
|
|
117
|
+
|
|
118
|
+
# Process the measurement table if exists
|
|
119
|
+
if MEASUREMENT in domain_table_list:
|
|
120
|
+
measurement = preprocess_domain_table(spark, input_folder, MEASUREMENT)
|
|
121
|
+
required_measurement = preprocess_domain_table(spark, input_folder, REQUIRED_MEASUREMENT)
|
|
122
|
+
# The select is necessary to make sure the order of the columns is the same as the
|
|
123
|
+
# original dataframe, otherwise the union might use the wrong columns
|
|
124
|
+
scaled_measurement = process_measurement(spark, measurement, required_measurement, output_folder)
|
|
125
|
+
|
|
126
|
+
if patient_events:
|
|
127
|
+
# Union all measurement records together with other domain records
|
|
128
|
+
patient_events = patient_events.unionByName(scaled_measurement)
|
|
129
|
+
else:
|
|
130
|
+
patient_events = scaled_measurement
|
|
131
|
+
|
|
132
|
+
patient_events = (
|
|
133
|
+
patient_events.join(visit_occurrence_person, "visit_occurrence_id")
|
|
134
|
+
.select(
|
|
135
|
+
[patient_events[fieldName] for fieldName in patient_events.schema.fieldNames()]
|
|
136
|
+
+ ["visit_concept_id", "age"]
|
|
137
|
+
)
|
|
138
|
+
.withColumn("cohort_member_id", F.col("person_id"))
|
|
139
|
+
)
|
|
140
|
+
|
|
141
|
+
# Apply the age security measure
|
|
142
|
+
# We only keep the patient records, whose corresponding age is less than 90
|
|
143
|
+
if apply_age_filter:
|
|
144
|
+
patient_events = patient_events.where(F.col("age") < 90)
|
|
145
|
+
|
|
146
|
+
if not continue_from_events:
|
|
147
|
+
patient_events.write.mode("overwrite").parquet(os.path.join(output_folder, "all_patient_events"))
|
|
148
|
+
|
|
149
|
+
patient_events = spark.read.parquet(os.path.join(output_folder, "all_patient_events"))
|
|
150
|
+
|
|
151
|
+
if is_new_patient_representation:
|
|
152
|
+
sequence_data = create_sequence_data_with_att(
|
|
153
|
+
patient_events,
|
|
154
|
+
visit_occurrence_person,
|
|
155
|
+
date_filter=date_filter,
|
|
156
|
+
include_visit_type=include_visit_type,
|
|
157
|
+
exclude_visit_tokens=exclude_visit_tokens,
|
|
158
|
+
patient_demographic=person if gpt_patient_sequence else None,
|
|
159
|
+
death=death,
|
|
160
|
+
att_type=att_type,
|
|
161
|
+
exclude_demographic=exclude_demographic,
|
|
162
|
+
use_age_group=use_age_group,
|
|
163
|
+
include_inpatient_hour_token=include_inpatient_hour_token,
|
|
164
|
+
)
|
|
165
|
+
else:
|
|
166
|
+
sequence_data = create_sequence_data(
|
|
167
|
+
patient_events,
|
|
168
|
+
date_filter=date_filter,
|
|
169
|
+
include_visit_type=include_visit_type,
|
|
170
|
+
classic_bert_seq=is_classic_bert,
|
|
171
|
+
)
|
|
172
|
+
|
|
173
|
+
if include_prolonged_stay:
|
|
174
|
+
udf = F.when(
|
|
175
|
+
F.col("visit_concept_id").isin([9201, 262, 9203]),
|
|
176
|
+
F.coalesce(
|
|
177
|
+
(F.datediff("visit_end_date", "visit_start_date") > 7).cast("int"),
|
|
178
|
+
F.lit(0),
|
|
179
|
+
),
|
|
180
|
+
).otherwise(F.lit(0))
|
|
181
|
+
visit_occurrence = preprocess_domain_table(spark, input_folder, VISIT_OCCURRENCE)
|
|
182
|
+
visit_occurrence = (
|
|
183
|
+
visit_occurrence.withColumn("prolonged_length_stay", udf)
|
|
184
|
+
.select("person_id", "prolonged_length_stay")
|
|
185
|
+
.withColumn(
|
|
186
|
+
"prolonged_length_stay",
|
|
187
|
+
F.max("prolonged_length_stay").over(Window.partitionBy("person_id")),
|
|
188
|
+
)
|
|
189
|
+
.distinct()
|
|
190
|
+
)
|
|
191
|
+
sequence_data = sequence_data.join(visit_occurrence, "person_id")
|
|
192
|
+
|
|
193
|
+
if include_sequence_information_content:
|
|
194
|
+
concept_df = patient_events.select("person_id", F.col("standard_concept_id").alias("concept_id"))
|
|
195
|
+
concept_freq = (
|
|
196
|
+
concept_df.groupBy("concept_id")
|
|
197
|
+
.count()
|
|
198
|
+
.withColumn("prob", F.col("count") / F.sum("count").over(Window.partitionBy()))
|
|
199
|
+
.withColumn("ic", -F.log("prob"))
|
|
200
|
+
)
|
|
201
|
+
|
|
202
|
+
patient_ic_df = concept_df.join(concept_freq, "concept_id").groupby("person_id").agg(F.mean("ic").alias("ic"))
|
|
203
|
+
|
|
204
|
+
sequence_data = sequence_data.join(patient_ic_df, "person_id")
|
|
205
|
+
|
|
206
|
+
patient_splits_folder = os.path.join(input_folder, "patient_splits")
|
|
207
|
+
if os.path.exists(patient_splits_folder):
|
|
208
|
+
patient_splits = spark.read.parquet(patient_splits_folder)
|
|
209
|
+
sequence_data.join(patient_splits, "person_id").write.mode("overwrite").parquet(
|
|
210
|
+
os.path.join(output_folder, "patient_sequence", "temp")
|
|
211
|
+
)
|
|
212
|
+
sequence_data = spark.read.parquet(os.path.join(output_folder, "patient_sequence", "temp"))
|
|
213
|
+
sequence_data.where('split="train"').write.mode("overwrite").parquet(
|
|
214
|
+
os.path.join(output_folder, "patient_sequence/train")
|
|
215
|
+
)
|
|
216
|
+
sequence_data.where('split="test"').write.mode("overwrite").parquet(
|
|
217
|
+
os.path.join(output_folder, "patient_sequence/test")
|
|
218
|
+
)
|
|
219
|
+
shutil.rmtree(os.path.join(output_folder, "patient_sequence", "temp"))
|
|
220
|
+
else:
|
|
221
|
+
sequence_data.write.mode("overwrite").parquet(os.path.join(output_folder, "patient_sequence"))
|
|
222
|
+
|
|
223
|
+
|
|
224
|
+
if __name__ == "__main__":
|
|
225
|
+
parser = argparse.ArgumentParser(description="Arguments for generate training data for Bert")
|
|
226
|
+
parser.add_argument(
|
|
227
|
+
"-i",
|
|
228
|
+
"--input_folder",
|
|
229
|
+
dest="input_folder",
|
|
230
|
+
action="store",
|
|
231
|
+
help="The path for your input_folder where the raw data is",
|
|
232
|
+
required=True,
|
|
233
|
+
)
|
|
234
|
+
parser.add_argument(
|
|
235
|
+
"-o",
|
|
236
|
+
"--output_folder",
|
|
237
|
+
dest="output_folder",
|
|
238
|
+
action="store",
|
|
239
|
+
help="The path for your output_folder",
|
|
240
|
+
required=True,
|
|
241
|
+
)
|
|
242
|
+
parser.add_argument(
|
|
243
|
+
"-tc",
|
|
244
|
+
"--domain_table_list",
|
|
245
|
+
dest="domain_table_list",
|
|
246
|
+
nargs="+",
|
|
247
|
+
action="store",
|
|
248
|
+
help="The list of domain tables you want to download",
|
|
249
|
+
type=validate_table_names,
|
|
250
|
+
required=True,
|
|
251
|
+
)
|
|
252
|
+
parser.add_argument(
|
|
253
|
+
"-d",
|
|
254
|
+
"--date_filter",
|
|
255
|
+
dest="date_filter",
|
|
256
|
+
type=lambda s: datetime.datetime.strptime(s, "%Y-%m-%d"),
|
|
257
|
+
action="store",
|
|
258
|
+
required=False,
|
|
259
|
+
default="2018-01-01",
|
|
260
|
+
)
|
|
261
|
+
parser.add_argument(
|
|
262
|
+
"-iv",
|
|
263
|
+
"--include_visit_type",
|
|
264
|
+
dest="include_visit_type",
|
|
265
|
+
action="store_true",
|
|
266
|
+
help="Specify whether to include visit types for generating the training data",
|
|
267
|
+
)
|
|
268
|
+
parser.add_argument(
|
|
269
|
+
"-ip",
|
|
270
|
+
"--is_new_patient_representation",
|
|
271
|
+
dest="is_new_patient_representation",
|
|
272
|
+
action="store_true",
|
|
273
|
+
help="Specify whether to generate the sequence of EHR records using the new patient " "representation",
|
|
274
|
+
)
|
|
275
|
+
parser.add_argument(
|
|
276
|
+
"-ib",
|
|
277
|
+
"--is_classic_bert_sequence",
|
|
278
|
+
dest="is_classic_bert_sequence",
|
|
279
|
+
action="store_true",
|
|
280
|
+
help="Specify whether to generate the sequence of EHR records using the classic BERT " "sequence",
|
|
281
|
+
)
|
|
282
|
+
parser.add_argument(
|
|
283
|
+
"-ev",
|
|
284
|
+
"--exclude_visit_tokens",
|
|
285
|
+
dest="exclude_visit_tokens",
|
|
286
|
+
action="store_true",
|
|
287
|
+
help="Specify whether or not to exclude the VS and VE tokens",
|
|
288
|
+
)
|
|
289
|
+
parser.add_argument(
|
|
290
|
+
"--include_prolonged_length_stay",
|
|
291
|
+
dest="include_prolonged_stay",
|
|
292
|
+
action="store_true",
|
|
293
|
+
help="Specify whether or not to include the data for the second learning objective for " "Med-BERT",
|
|
294
|
+
)
|
|
295
|
+
parser.add_argument("--include_concept_list", dest="include_concept_list", action="store_true")
|
|
296
|
+
parser.add_argument("--gpt_patient_sequence", dest="gpt_patient_sequence", action="store_true")
|
|
297
|
+
parser.add_argument("--apply_age_filter", dest="apply_age_filter", action="store_true")
|
|
298
|
+
parser.add_argument("--include_death", dest="include_death", action="store_true")
|
|
299
|
+
parser.add_argument("--exclude_demographic", dest="exclude_demographic", action="store_true")
|
|
300
|
+
parser.add_argument("--use_age_group", dest="use_age_group", action="store_true")
|
|
301
|
+
parser.add_argument("--with_drug_rollup", dest="with_drug_rollup", action="store_true")
|
|
302
|
+
parser.add_argument(
|
|
303
|
+
"--include_inpatient_hour_token",
|
|
304
|
+
dest="include_inpatient_hour_token",
|
|
305
|
+
action="store_true",
|
|
306
|
+
)
|
|
307
|
+
parser.add_argument("--continue_from_events", dest="continue_from_events", action="store_true")
|
|
308
|
+
parser.add_argument(
|
|
309
|
+
"--att_type",
|
|
310
|
+
dest="att_type",
|
|
311
|
+
action="store",
|
|
312
|
+
choices=[e.value for e in AttType],
|
|
313
|
+
)
|
|
314
|
+
|
|
315
|
+
ARGS = parser.parse_args()
|
|
316
|
+
|
|
317
|
+
main(
|
|
318
|
+
ARGS.input_folder,
|
|
319
|
+
ARGS.output_folder,
|
|
320
|
+
ARGS.domain_table_list,
|
|
321
|
+
ARGS.date_filter,
|
|
322
|
+
ARGS.include_visit_type,
|
|
323
|
+
ARGS.is_new_patient_representation,
|
|
324
|
+
ARGS.exclude_visit_tokens,
|
|
325
|
+
ARGS.is_classic_bert_sequence,
|
|
326
|
+
ARGS.include_prolonged_stay,
|
|
327
|
+
ARGS.include_concept_list,
|
|
328
|
+
ARGS.gpt_patient_sequence,
|
|
329
|
+
ARGS.apply_age_filter,
|
|
330
|
+
ARGS.include_death,
|
|
331
|
+
AttType(ARGS.att_type),
|
|
332
|
+
exclude_demographic=ARGS.exclude_demographic,
|
|
333
|
+
use_age_group=ARGS.use_age_group,
|
|
334
|
+
with_drug_rollup=ARGS.with_drug_rollup,
|
|
335
|
+
include_inpatient_hour_token=ARGS.include_inpatient_hour_token,
|
|
336
|
+
continue_from_events=ARGS.continue_from_events,
|
|
337
|
+
)
|
|
File without changes
|