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
__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
|
+
)
|