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.
Files changed (57) hide show
  1. __init__.py +0 -0
  2. cehrbert_data/__init__.py +0 -0
  3. cehrbert_data/apps/__init__.py +0 -0
  4. cehrbert_data/apps/generate_concept_similarity_table.py +423 -0
  5. cehrbert_data/apps/generate_hierarchical_bert_training_data.py +238 -0
  6. cehrbert_data/apps/generate_included_concept_list.py +116 -0
  7. cehrbert_data/apps/generate_information_content.py +131 -0
  8. cehrbert_data/apps/generate_required_labs.py +112 -0
  9. cehrbert_data/apps/generate_training_data.py +337 -0
  10. cehrbert_data/cohorts/__init__.py +0 -0
  11. cehrbert_data/cohorts/atrial_fibrillation.py +45 -0
  12. cehrbert_data/cohorts/cabg.py +72 -0
  13. cehrbert_data/cohorts/coronary_artery_disease.py +84 -0
  14. cehrbert_data/cohorts/covid.py +43 -0
  15. cehrbert_data/cohorts/covid_inpatient.py +84 -0
  16. cehrbert_data/cohorts/death.py +46 -0
  17. cehrbert_data/cohorts/heart_failure.py +423 -0
  18. cehrbert_data/cohorts/ischemic_stroke.py +45 -0
  19. cehrbert_data/cohorts/last_visit_discharged_home.py +35 -0
  20. cehrbert_data/cohorts/query_builder.py +153 -0
  21. cehrbert_data/cohorts/spark_app_base.py +745 -0
  22. cehrbert_data/cohorts/type_two_diabietes.py +166 -0
  23. cehrbert_data/cohorts/ventilation.py +22 -0
  24. cehrbert_data/config/__init__.py +0 -0
  25. cehrbert_data/config/output_names.py +9 -0
  26. cehrbert_data/const/__init__.py +0 -0
  27. cehrbert_data/const/__pycache__/__init__.cpython-311.pyc +0 -0
  28. cehrbert_data/const/__pycache__/common.cpython-311.pyc +0 -0
  29. cehrbert_data/const/common.py +28 -0
  30. cehrbert_data/decorators/__init__.py +0 -0
  31. cehrbert_data/decorators/__pycache__/__init__.cpython-311.pyc +0 -0
  32. cehrbert_data/decorators/__pycache__/patient_event_decorator.cpython-311.pyc +0 -0
  33. cehrbert_data/decorators/patient_event_decorator.py +759 -0
  34. cehrbert_data/prediction_cohorts/__init__.py +0 -0
  35. cehrbert_data/prediction_cohorts/afib_ischemic_stroke.py +14 -0
  36. cehrbert_data/prediction_cohorts/cad_cabg_cohort.py +19 -0
  37. cehrbert_data/prediction_cohorts/cad_hf_cohort.py +14 -0
  38. cehrbert_data/prediction_cohorts/copd_readmission.py +76 -0
  39. cehrbert_data/prediction_cohorts/covid_death.py +14 -0
  40. cehrbert_data/prediction_cohorts/covid_ventilation.py +14 -0
  41. cehrbert_data/prediction_cohorts/discharge_home_death.py +18 -0
  42. cehrbert_data/prediction_cohorts/hf_readmission.py +82 -0
  43. cehrbert_data/prediction_cohorts/hospitalization.py +100 -0
  44. cehrbert_data/prediction_cohorts/hospitalization_mortality.py +77 -0
  45. cehrbert_data/prediction_cohorts/t2dm_hf_cohort.py +14 -0
  46. cehrbert_data/queries/__init__.py +0 -0
  47. cehrbert_data/queries/measurement_unit_stats_query.py +42 -0
  48. cehrbert_data/tools/__init__.py +0 -0
  49. cehrbert_data/tools/download_omop_tables.py +141 -0
  50. cehrbert_data/utils/__init__.py +0 -0
  51. cehrbert_data/utils/spark_parse_args.py +350 -0
  52. cehrbert_data/utils/spark_utils.py +1405 -0
  53. cehrbert_data-0.0.1.dist-info/LICENSE +21 -0
  54. cehrbert_data-0.0.1.dist-info/METADATA +116 -0
  55. cehrbert_data-0.0.1.dist-info/RECORD +57 -0
  56. cehrbert_data-0.0.1.dist-info/WHEEL +5 -0
  57. 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