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
__init__.py ADDED
File without changes
File without changes
File without changes
@@ -0,0 +1,423 @@
1
+ """This module provides functionality to extract patient event data from domain tables,.
2
+
3
+ compute information content and semantic similarity for concepts, and calculate concept
4
+ similarity scores.
5
+
6
+ Functions: extract_data: Extract data from specified domain tables. compute_information_content:
7
+ Compute the information content for concepts based on frequency.
8
+ compute_information_content_similarity: Compute the similarity between concepts based on
9
+ information content. compute_semantic_similarity: Compute the semantic similarity between concept
10
+ pairs. main: Main function to orchestrate the extraction, processing, and saving of concept
11
+ similarity data.
12
+ """
13
+
14
+ import datetime
15
+ import logging
16
+ import os
17
+ from typing import List
18
+
19
+ from pyspark.sql import DataFrame, SparkSession
20
+ from pyspark.sql import functions as F
21
+
22
+ from cehrbert_data.config.output_names import CONCEPT_SIMILARITY_PATH, QUALIFIED_CONCEPT_LIST_PATH
23
+ from cehrbert_data.const.common import CONCEPT, CONCEPT_ANCESTOR
24
+ from cehrbert_data.utils.spark_utils import join_domain_tables, preprocess_domain_table, validate_table_names
25
+
26
+
27
+ def extract_data(spark: SparkSession, input_folder: str, domain_table_list: List[str]):
28
+ """
29
+ Extract patient event data from the specified domain tables.
30
+
31
+ Args:
32
+ spark (SparkSession): The Spark session to use for processing.
33
+ input_folder (str): Path to the input folder containing domain tables.
34
+ domain_table_list (List[str]): List of domain table names to extract data from.
35
+
36
+ Returns:
37
+ DataFrame: A DataFrame containing extracted and processed patient event data.
38
+ """
39
+ domain_tables = []
40
+ for domain_table_name in domain_table_list:
41
+ domain_tables.append(preprocess_domain_table(spark, input_folder, domain_table_name))
42
+ patient_event = join_domain_tables(domain_tables)
43
+ # Remove all concept_id records
44
+ patient_event = patient_event.where("standard_concept_id <> 0")
45
+
46
+ return patient_event
47
+
48
+
49
+ def compute_information_content(patient_event: DataFrame, concept_ancestor: DataFrame):
50
+ """
51
+ Calculate the information content using the frequency of each concept and the graph.
52
+
53
+ :param patient_event:
54
+ :param concept_ancestor:
55
+ :return:
56
+ """
57
+ # Get the total count
58
+ total_count = patient_event.distinct().count()
59
+ # Count the frequency of each concept
60
+ concept_frequency = patient_event.distinct().groupBy("standard_concept_id").count()
61
+ # left join b/w descendent_concept_id and the standard_concept_id in the concept freq table
62
+ freq_df = (
63
+ concept_frequency.join(
64
+ concept_ancestor,
65
+ F.col("descendant_concept_id") == F.col("standard_concept_id"),
66
+ )
67
+ .groupBy("ancestor_concept_id")
68
+ .sum("count")
69
+ .withColumnRenamed("ancestor_concept_id", "concept_id")
70
+ .withColumnRenamed("sum(count)", "count")
71
+ )
72
+ # Calculate information content for each concept
73
+ information_content = freq_df.withColumn("information_content", (-F.log(F.col("count") / total_count))).withColumn(
74
+ "probability", F.col("count") / total_count
75
+ )
76
+
77
+ return information_content
78
+
79
+
80
+ def compute_information_content_similarity(
81
+ concept_pair: DataFrame, information_content: DataFrame, concept_ancestor: DataFrame
82
+ ):
83
+ """
84
+ Compute the similarity between concept pairs based on their information content.
85
+
86
+ Args:
87
+ concept_pair (DataFrame): A DataFrame containing pairs of concepts.
88
+ information_content (DataFrame): A DataFrame with information content for concepts.
89
+ concept_ancestor (DataFrame): A DataFrame containing concept ancestor relationships.
90
+
91
+ Returns:
92
+ DataFrame: A DataFrame containing various similarity measures for concept pairs.
93
+ """
94
+ # Extract the pairs of concepts from the training data and join to the information content table
95
+ information_content_concept_pair = (
96
+ concept_pair.select("concept_id_1", "concept_id_2")
97
+ .join(
98
+ information_content,
99
+ F.col("concept_id_1") == F.col("concept_id"),
100
+ "left_outer",
101
+ )
102
+ .select(
103
+ F.col("concept_id_1"),
104
+ F.col("concept_id_2"),
105
+ F.col("information_content").alias("information_content_1"),
106
+ )
107
+ .join(
108
+ information_content,
109
+ F.col("concept_id_2") == F.col("concept_id"),
110
+ "left_outer",
111
+ )
112
+ .select(
113
+ F.col("concept_id_1"),
114
+ F.col("concept_id_2"),
115
+ F.col("information_content_1"),
116
+ F.col("information_content").alias("information_content_2"),
117
+ )
118
+ )
119
+
120
+ # Join to get all the ancestors of concept_id_1
121
+ concept_id_1_ancestor = information_content_concept_pair.join(
122
+ concept_ancestor, F.col("concept_id_1") == F.col("descendant_concept_id")
123
+ ).select("concept_id_1", "concept_id_2", "ancestor_concept_id")
124
+
125
+ # Join to get all the ancestors of concept_id_2
126
+ concept_id_2_ancestor = concept_pair.join(
127
+ concept_ancestor, F.col("concept_id_2") == F.col("descendant_concept_id")
128
+ ).select("concept_id_1", "concept_id_2", "ancestor_concept_id")
129
+
130
+ # Compute the summed information content of all ancestors of concept_id_1 and concept_id_2
131
+ union_sum = (
132
+ concept_id_1_ancestor.union(concept_id_2_ancestor)
133
+ .distinct()
134
+ .join(information_content, F.col("ancestor_concept_id") == F.col("concept_id"))
135
+ .groupBy("concept_id_1", "concept_id_2")
136
+ .agg(F.sum("information_content").alias("ancestor_union_ic"))
137
+ )
138
+
139
+ # Compute the summed information content of common ancestors of concept_id_1 and concept_id_2
140
+ intersection_sum = (
141
+ concept_id_1_ancestor.intersect(concept_id_2_ancestor)
142
+ .join(information_content, F.col("ancestor_concept_id") == F.col("concept_id"))
143
+ .groupBy("concept_id_1", "concept_id_2")
144
+ .agg(F.sum("information_content").alias("ancestor_intersection_ic"))
145
+ )
146
+
147
+ # Compute the information content and probability of the most informative common ancestor (MICA)
148
+ mica_ancestor = (
149
+ concept_id_1_ancestor.intersect(concept_id_2_ancestor)
150
+ .join(information_content, F.col("ancestor_concept_id") == F.col("concept_id"))
151
+ .groupBy("concept_id_1", "concept_id_2")
152
+ .agg(
153
+ F.max("information_content").alias("mica_information_content"),
154
+ F.max("probability").alias("mica_probability"),
155
+ )
156
+ )
157
+
158
+ # Join the MICA to pairs of concepts
159
+ features = information_content_concept_pair.join(
160
+ mica_ancestor,
161
+ (information_content_concept_pair["concept_id_1"] == mica_ancestor["concept_id_1"])
162
+ & (information_content_concept_pair["concept_id_2"] == mica_ancestor["concept_id_2"]),
163
+ "left_outer",
164
+ ).select(
165
+ [information_content_concept_pair[f] for f in information_content_concept_pair.schema.fieldNames()]
166
+ + [F.col("mica_information_content"), F.col("mica_probability")]
167
+ )
168
+
169
+ # Compute the lin measure
170
+ features = features.withColumn(
171
+ "lin_measure",
172
+ 2 * F.col("mica_information_content") / (F.col("information_content_1") * F.col("information_content_2")),
173
+ )
174
+
175
+ # Compute the jiang measure
176
+ features = features.withColumn(
177
+ "jiang_measure",
178
+ 1 - (F.col("information_content_1") + F.col("information_content_2") - 2 * F.col("mica_information_content")),
179
+ )
180
+
181
+ # Compute the information coefficient
182
+ features = features.withColumn(
183
+ "information_coefficient",
184
+ F.col("lin_measure") * (1 - 1 / (1 + F.col("mica_information_content"))),
185
+ )
186
+
187
+ # Compute the relevance_measure
188
+ features = features.withColumn("relevance_measure", F.col("lin_measure") * (1 - F.col("mica_probability")))
189
+
190
+ # Join to get the summed information content of the common ancestors of concept_id_1 and
191
+ # concept_id_2
192
+ features = features.join(
193
+ intersection_sum,
194
+ (features["concept_id_1"] == intersection_sum["concept_id_1"])
195
+ & (features["concept_id_2"] == intersection_sum["concept_id_2"]),
196
+ "left_outer",
197
+ ).select([features[f] for f in features.schema.fieldNames()] + [F.col("ancestor_intersection_ic")])
198
+
199
+ # Join to get the summed information content of the common ancestors of concept_id_1 and
200
+ # concept_id_2
201
+ features = features.join(
202
+ union_sum,
203
+ (features["concept_id_1"] == union_sum["concept_id_1"])
204
+ & (features["concept_id_2"] == union_sum["concept_id_2"]),
205
+ "left_outer",
206
+ ).select([features[f] for f in features.schema.fieldNames()] + [F.col("ancestor_union_ic")])
207
+
208
+ # Compute the graph information content measure
209
+ features = features.withColumn(
210
+ "graph_ic_measure",
211
+ F.col("ancestor_intersection_ic") / F.col("ancestor_union_ic"),
212
+ )
213
+
214
+ return features.select(
215
+ [
216
+ F.col("concept_id_1"),
217
+ F.col("concept_id_2"),
218
+ F.col("mica_information_content"),
219
+ F.col("lin_measure"),
220
+ F.col("jiang_measure"),
221
+ F.col("information_coefficient"),
222
+ F.col("relevance_measure"),
223
+ F.col("graph_ic_measure"),
224
+ ]
225
+ )
226
+
227
+
228
+ def compute_semantic_similarity(spark, patient_event, concept, concept_ancestor):
229
+ required_concept = (
230
+ patient_event.distinct()
231
+ .select("standard_concept_id")
232
+ .join(concept, F.col("standard_concept_id") == F.col("concept_id"))
233
+ .select("standard_concept_id", "domain_id")
234
+ )
235
+
236
+ concept_ancestor.createOrReplaceTempView("concept_ancestor")
237
+ required_concept.createOrReplaceTempView("required_concept")
238
+
239
+ concept_pair = spark.sql(
240
+ """
241
+ WITH concept_pair AS (
242
+ SELECT
243
+ c1.standard_concept_id AS concept_id_1,
244
+ c2.standard_concept_id AS concept_id_2,
245
+ c1.domain_id
246
+ FROM required_concept AS c1
247
+ JOIN required_concept AS c2
248
+ ON c1.domain_id = c2.domain_id
249
+ WHERE c1.standard_concept_id <> c2.standard_concept_id
250
+ )
251
+ SELECT
252
+ cp.concept_id_1,
253
+ cp.concept_id_2,
254
+ ca_1.ancestor_concept_id AS common_ancestor_concept_id,
255
+ ca_1.min_levels_of_separation AS distance_1,
256
+ ca_2.min_levels_of_separation AS distance_2
257
+ FROM concept_pair AS cp
258
+ JOIN concept_ancestor AS ca_1
259
+ ON cp.concept_id_1 = ca_1.descendant_concept_id
260
+ JOIN concept_ancestor AS ca_2
261
+ ON cp.concept_id_2 = ca_2.descendant_concept_id
262
+ WHERE ca_1.ancestor_concept_id = ca_2.ancestor_concept_id
263
+ """
264
+ )
265
+
266
+ # Find the root concepts
267
+ root_concept = (
268
+ concept_ancestor.groupBy("descendant_concept_id")
269
+ .count()
270
+ .where("count = 1")
271
+ .withColumnRenamed("descendant_concept_id", "root_concept_id")
272
+ )
273
+ # Retrieve all ancestor descendant relationships for the root concepts
274
+ root_concept_relationship = (
275
+ root_concept.join(
276
+ concept_ancestor,
277
+ root_concept["root_concept_id"] == concept_ancestor["ancestor_concept_id"],
278
+ )
279
+ .select(
280
+ concept_ancestor["ancestor_concept_id"],
281
+ concept_ancestor["descendant_concept_id"],
282
+ concept_ancestor["max_levels_of_separation"].alias("root_distance"),
283
+ )
284
+ .where("ancestor_concept_id <> descendant_concept_id")
285
+ )
286
+
287
+ # Join to get all root concepts and their corresponding root_distance
288
+ concept_pair = concept_pair.join(
289
+ root_concept_relationship,
290
+ F.col("common_ancestor_concept_id") == F.col("descendant_concept_id"),
291
+ ).select("concept_id_1", "concept_id_2", "distance_1", "distance_2", "root_distance")
292
+
293
+ # Compute the semantic similarity
294
+ concept_pair_similarity = concept_pair.withColumn(
295
+ "semantic_similarity",
296
+ 2 * F.col("root_distance") / (2 * F.col("root_distance") + F.col("distance_1") + F.col("distance_2")),
297
+ )
298
+ # Find the maximum semantic similarity
299
+ concept_pair_similarity = concept_pair_similarity.groupBy("concept_id_1", "concept_id_2").agg(
300
+ F.max("semantic_similarity").alias("semantic_similarity")
301
+ )
302
+
303
+ return concept_pair_similarity
304
+
305
+
306
+ def main(
307
+ input_folder: str,
308
+ output_folder: str,
309
+ domain_table_list: List[str],
310
+ date_filter: str,
311
+ include_concept_list: bool,
312
+ ):
313
+ """
314
+ Main function to generate the concept similarity table.
315
+
316
+ Args:
317
+ input_folder (str): The path to the input folder containing raw data.
318
+ output_folder (str): The path to the output folder to store the results.
319
+ domain_table_list (List[str]): List of domain tables to process.
320
+ date_filter (str): Date filter to apply to the data.
321
+ include_concept_list (bool): Whether to include a filtered concept list.
322
+ """
323
+
324
+ spark = SparkSession.builder.appName("Generate the concept similarity table").getOrCreate()
325
+
326
+ logger = logging.getLogger(__name__)
327
+ logger.info(
328
+ "input_folder: %s\noutput_folder: %s\ndomain_table_list: %s\ndate_filter: " "%s\ninclude_concept_list: %s",
329
+ input_folder,
330
+ output_folder,
331
+ domain_table_list,
332
+ date_filter,
333
+ include_concept_list,
334
+ )
335
+
336
+ concept = preprocess_domain_table(spark, input_folder, CONCEPT)
337
+ concept_ancestor = preprocess_domain_table(spark, input_folder, CONCEPT_ANCESTOR)
338
+
339
+ # Extract all data points from specified domains
340
+ patient_event = extract_data(spark, input_folder, domain_table_list)
341
+
342
+ # Calculate information content using unfiltered the patient event dataframe
343
+ information_content = compute_information_content(patient_event, concept_ancestor)
344
+
345
+ # Filter out concepts that are not required in the required concept_list
346
+ if include_concept_list and patient_event:
347
+ # Filter out concepts
348
+ qualified_concepts = F.broadcast(preprocess_domain_table(spark, input_folder, QUALIFIED_CONCEPT_LIST_PATH))
349
+
350
+ patient_event = patient_event.join(qualified_concepts, "standard_concept_id").select("standard_concept_id")
351
+
352
+ concept_pair_similarity = compute_semantic_similarity(spark, patient_event, concept, concept_ancestor)
353
+
354
+ # Compute the information content based similarity scores
355
+ concept_pair_ic_similarity = compute_information_content_similarity(
356
+ concept_pair_similarity, information_content, concept_ancestor
357
+ )
358
+
359
+ concept_pair_similarity_columns = [concept_pair_similarity[f] for f in concept_pair_similarity.schema.fieldNames()]
360
+ concept_pair_ic_similarity_columns = [
361
+ f for f in concept_pair_ic_similarity.schema.fieldNames() if "concept_id" not in f
362
+ ]
363
+
364
+ # Join two dataframes to get the final result
365
+ concept_pair_similarity = concept_pair_similarity.join(
366
+ concept_pair_ic_similarity,
367
+ (concept_pair_similarity["concept_id_1"] == concept_pair_ic_similarity["concept_id_1"])
368
+ & (concept_pair_similarity["concept_id_2"] == concept_pair_ic_similarity["concept_id_2"]),
369
+ ).select(concept_pair_similarity_columns + concept_pair_ic_similarity_columns)
370
+
371
+ concept_pair_similarity.write.mode("overwrite").parquet(os.path.join(output_folder, CONCEPT_SIMILARITY_PATH))
372
+
373
+
374
+ if __name__ == "__main__":
375
+ import argparse
376
+
377
+ parser = argparse.ArgumentParser(description="Arguments for generate Concept Similarity Table")
378
+ parser.add_argument(
379
+ "-i",
380
+ "--input_folder",
381
+ dest="input_folder",
382
+ action="store",
383
+ help="The path for your input_folder where the raw data is",
384
+ required=True,
385
+ )
386
+ parser.add_argument(
387
+ "-o",
388
+ "--output_folder",
389
+ dest="output_folder",
390
+ action="store",
391
+ help="The path for your output_folder",
392
+ required=True,
393
+ )
394
+ parser.add_argument(
395
+ "-tc",
396
+ "--domain_table_list",
397
+ dest="domain_table_list",
398
+ nargs="+",
399
+ action="store",
400
+ help="The list of domain tables you want to download",
401
+ type=validate_table_names,
402
+ required=True,
403
+ )
404
+ parser.add_argument(
405
+ "-d",
406
+ "--date_filter",
407
+ dest="date_filter",
408
+ type=lambda s: datetime.datetime.strptime(s, "%Y-%m-%d"),
409
+ action="store",
410
+ required=False,
411
+ default="2018-01-01",
412
+ )
413
+ parser.add_argument("--include_concept_list", dest="include_concept_list", action="store_true")
414
+
415
+ ARGS = parser.parse_args()
416
+
417
+ main(
418
+ ARGS.input_folder,
419
+ ARGS.output_folder,
420
+ ARGS.domain_table_list,
421
+ ARGS.date_filter,
422
+ ARGS.include_concept_list,
423
+ )
@@ -0,0 +1,238 @@
1
+ """
2
+ This module generates hierarchical BERT training data based on domain tables from OMOP EHR data.
3
+
4
+ It processes patient event data, joins multiple domain tables, filters concepts based on a
5
+ minimum number of patients, and creates hierarchical sequence data for BERT training.
6
+
7
+ Key Functions:
8
+ - preprocess_domain_table: Preprocesses domain tables for data extraction.
9
+ - process_measurement: Handles special processing for measurement data.
10
+ - join_domain_tables: Joins multiple domain tables into a unified DataFrame.
11
+ - create_hierarchical_sequence_data: Generates hierarchical sequence data for training.
12
+
13
+ Command-line Arguments:
14
+ - input_folder: Path to the directory containing input data.
15
+ - output_folder: Path to the directory where the output will be saved.
16
+ - domain_table_list: List of domain tables to process.
17
+ - date_filter: Optional filter for processing the data based on date.
18
+ - max_num_of_visits_per_person: Maximum number of visits per patient to include.
19
+ - min_observation_period: Minimum observation period in days for patients to be included.
20
+ - include_concept_list: Whether to apply a filter to retain certain concepts.
21
+ - include_incomplete_visit: Whether to include incomplete visit records in the training data.
22
+ """
23
+
24
+ import datetime
25
+ import logging
26
+ import os
27
+
28
+ from pyspark.sql import SparkSession
29
+ from pyspark.sql import functions as F
30
+
31
+ from cehrbert_data.config.output_names import PARQUET_DATA_PATH, QUALIFIED_CONCEPT_LIST_PATH
32
+ from cehrbert_data.const.common import MEASUREMENT, OBSERVATION_PERIOD, PERSON, REQUIRED_MEASUREMENT, VISIT_OCCURRENCE
33
+ from cehrbert_data.utils.spark_utils import (
34
+ create_hierarchical_sequence_data,
35
+ join_domain_tables,
36
+ preprocess_domain_table,
37
+ process_measurement,
38
+ validate_table_names,
39
+ )
40
+
41
+
42
+ def main(
43
+ input_folder,
44
+ output_folder,
45
+ domain_table_list,
46
+ date_filter,
47
+ max_num_of_visits_per_person,
48
+ min_observation_period: int = 360,
49
+ include_concept_list: bool = True,
50
+ include_incomplete_visit: bool = True,
51
+ ):
52
+ """
53
+ Main function to generate hierarchical BERT training data from domain tables.
54
+
55
+ Args:
56
+ input_folder (str): The path to the input folder containing raw data.
57
+ output_folder (str): The path to the output folder for storing the training data.
58
+ domain_table_list (list): A list of domain tables to process (e.g., condition_occurrence).
59
+ date_filter (str): Date filter for processing data, default is '2018-01-01'.
60
+ max_num_of_visits_per_person (int): The maximum number of visits to include per person.
61
+ min_observation_period (int, optional): Minimum observation period in days. Default is 360.
62
+ include_concept_list (bool, optional): Whether to filter by concept list. Default is True.
63
+ include_incomplete_visit (bool, optional): Whether to include incomplete visits. Default is
64
+ True.
65
+
66
+ This function preprocesses domain tables, filters and processes measurement data,
67
+ and generates hierarchical sequence data for training BERT models on EHR records.
68
+ """
69
+ spark = SparkSession.builder.appName("Generate Hierarchical Bert Training Data").getOrCreate()
70
+
71
+ logger = logging.getLogger(__name__)
72
+ logger.info(
73
+ "input_folder: %s\n"
74
+ "output_folder: %s\n"
75
+ "domain_table_list: %s\n"
76
+ "date_filter: %s\n"
77
+ "max_num_of_visits_per_person: %s\n"
78
+ "min_observation_period: %s\n"
79
+ "include_concept_list: %s\n"
80
+ "include_incomplete_visit: %s",
81
+ input_folder,
82
+ output_folder,
83
+ domain_table_list,
84
+ date_filter,
85
+ max_num_of_visits_per_person,
86
+ min_observation_period,
87
+ include_concept_list,
88
+ include_incomplete_visit,
89
+ )
90
+
91
+ domain_tables = []
92
+ # Exclude measurement from domain_table_list if exists because we need to process measurement
93
+ # in a different way
94
+ for domain_table_name in domain_table_list:
95
+ if domain_table_name != MEASUREMENT:
96
+ domain_tables.append(preprocess_domain_table(spark, input_folder, domain_table_name))
97
+
98
+ observation_period = (
99
+ preprocess_domain_table(spark, input_folder, OBSERVATION_PERIOD)
100
+ .withColumn(
101
+ "observation_period_start_date",
102
+ F.col("observation_period_start_date").cast("date"),
103
+ )
104
+ .withColumn(
105
+ "observation_period_end_date",
106
+ F.col("observation_period_end_date").cast("date"),
107
+ )
108
+ .withColumn(
109
+ "period",
110
+ F.datediff("observation_period_end_date", "observation_period_start_date"),
111
+ )
112
+ .where(F.col("period") >= min_observation_period)
113
+ .select("person_id")
114
+ )
115
+
116
+ visit_occurrence = preprocess_domain_table(spark, input_folder, VISIT_OCCURRENCE)
117
+ person = preprocess_domain_table(spark, input_folder, PERSON)
118
+
119
+ # Filter for the persons that have enough observation period
120
+ person = person.join(observation_period, "person_id").select([person[f] for f in person.schema.fieldNames()])
121
+
122
+ # Union all domain table records
123
+ patient_events = join_domain_tables(domain_tables)
124
+
125
+ column_names = patient_events.schema.fieldNames()
126
+
127
+ if include_concept_list and patient_events:
128
+ # Filter out concepts
129
+ qualified_concepts = F.broadcast(preprocess_domain_table(spark, input_folder, QUALIFIED_CONCEPT_LIST_PATH))
130
+ # The select is necessary to make sure the order of the columns is the same as the
131
+ # original dataframe
132
+ patient_events = patient_events.join(qualified_concepts, "standard_concept_id").select(column_names)
133
+
134
+ # Process the measurement table if exists
135
+ if MEASUREMENT in domain_table_list:
136
+ measurement = preprocess_domain_table(spark, input_folder, MEASUREMENT)
137
+ required_measurement = preprocess_domain_table(spark, input_folder, REQUIRED_MEASUREMENT)
138
+ # The select is necessary to make sure the order of the columns is the same as the
139
+ # original dataframe, otherwise the union might use the wrong columns
140
+ scaled_measurement = process_measurement(spark, measurement, required_measurement).select(column_names)
141
+
142
+ if patient_events:
143
+ # Union all measurement records together with other domain records
144
+ patient_events = patient_events.union(scaled_measurement)
145
+ else:
146
+ patient_events = scaled_measurement
147
+
148
+ # cohort_member_id is the same as the person_id
149
+ patient_events = patient_events.withColumn("cohort_member_id", F.col("person_id"))
150
+
151
+ sequence_data = create_hierarchical_sequence_data(
152
+ person,
153
+ visit_occurrence,
154
+ patient_events,
155
+ date_filter=date_filter,
156
+ max_num_of_visits_per_person=max_num_of_visits_per_person,
157
+ include_incomplete_visit=include_incomplete_visit,
158
+ )
159
+
160
+ sequence_data.write.mode("overwrite").parquet(os.path.join(output_folder, PARQUET_DATA_PATH))
161
+
162
+
163
+ if __name__ == "__main__":
164
+ import argparse
165
+
166
+ parser = argparse.ArgumentParser(description="Arguments for generate training data for Hierarchical Bert")
167
+ parser.add_argument(
168
+ "-i",
169
+ "--input_folder",
170
+ dest="input_folder",
171
+ action="store",
172
+ help="The path for your input_folder where the raw data is",
173
+ required=True,
174
+ )
175
+ parser.add_argument(
176
+ "-o",
177
+ "--output_folder",
178
+ dest="output_folder",
179
+ action="store",
180
+ help="The path for your output_folder",
181
+ required=True,
182
+ )
183
+ parser.add_argument(
184
+ "-tc",
185
+ "--domain_table_list",
186
+ dest="domain_table_list",
187
+ nargs="+",
188
+ action="store",
189
+ help="The list of domain tables you want to download",
190
+ type=validate_table_names,
191
+ required=True,
192
+ )
193
+ parser.add_argument(
194
+ "-d",
195
+ "--date_filter",
196
+ dest="date_filter",
197
+ type=lambda s: datetime.datetime.strptime(s, "%Y-%m-%d"),
198
+ action="store",
199
+ required=False,
200
+ default="2018-01-01",
201
+ )
202
+ parser.add_argument(
203
+ "--max_num_of_visits",
204
+ dest="max_num_of_visits",
205
+ action="store",
206
+ type=int,
207
+ default=200,
208
+ help="Max no.of visits per patient to be included",
209
+ required=False,
210
+ )
211
+ parser.add_argument(
212
+ "--min_observation_period",
213
+ dest="min_observation_period",
214
+ action="store",
215
+ type=int,
216
+ default=1,
217
+ help="Minimum observation period in days",
218
+ required=False,
219
+ )
220
+ parser.add_argument("--include_concept_list", dest="include_concept_list", action="store_true")
221
+ parser.add_argument(
222
+ "--include_incomplete_visit",
223
+ dest="include_incomplete_visit",
224
+ action="store_true",
225
+ )
226
+
227
+ ARGS = parser.parse_args()
228
+
229
+ main(
230
+ input_folder=ARGS.input_folder,
231
+ output_folder=ARGS.output_folder,
232
+ domain_table_list=ARGS.domain_table_list,
233
+ date_filter=ARGS.date_filter,
234
+ max_num_of_visits_per_person=ARGS.max_num_of_visits,
235
+ min_observation_period=ARGS.min_observation_period,
236
+ include_concept_list=ARGS.include_concept_list,
237
+ include_incomplete_visit=ARGS.include_incomplete_visit,
238
+ )