cehrbert-data 0.0.2__tar.gz → 0.0.4__tar.gz

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 (137) hide show
  1. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/PKG-INFO +1 -1
  2. cehrbert_data-0.0.4/src/cehrbert_data/apps/generate_required_labs.py +140 -0
  3. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/apps/generate_training_data.py +60 -36
  4. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/cohorts/spark_app_base.py +97 -80
  5. cehrbert_data-0.0.4/src/cehrbert_data/const/artificial_tokens.py +3 -0
  6. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/const/common.py +5 -0
  7. cehrbert_data-0.0.4/src/cehrbert_data/decorators/__init__.py +5 -0
  8. cehrbert_data-0.0.4/src/cehrbert_data/decorators/artificial_time_token_decorator.py +336 -0
  9. cehrbert_data-0.0.4/src/cehrbert_data/decorators/clinical_event_decorator.py +168 -0
  10. cehrbert_data-0.0.4/src/cehrbert_data/decorators/death_event_decorator.py +114 -0
  11. cehrbert_data-0.0.4/src/cehrbert_data/decorators/demographic_event_decorator.py +111 -0
  12. cehrbert_data-0.0.4/src/cehrbert_data/decorators/patient_event_decorator_base.py +146 -0
  13. cehrbert_data-0.0.4/src/cehrbert_data/decorators/token_priority.py +22 -0
  14. cehrbert_data-0.0.2/src/cehrbert_data/queries/measurement_unit_stats_query.py → cehrbert_data-0.0.4/src/cehrbert_data/queries/measurement_queries.py +17 -2
  15. cehrbert_data-0.0.4/src/cehrbert_data/utils/logging_utils.py +12 -0
  16. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/utils/spark_parse_args.py +27 -3
  17. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/utils/spark_utils.py +224 -133
  18. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data.egg-info/PKG-INFO +1 -1
  19. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data.egg-info/SOURCES.txt +11 -5
  20. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/tests/integration_tests/test_generate_training_data.py +2 -2
  21. cehrbert_data-0.0.4/tests/unit_tests/test_spark_utils.py +53 -0
  22. cehrbert_data-0.0.2/src/cehrbert_data/apps/generate_required_labs.py +0 -112
  23. cehrbert_data-0.0.2/src/cehrbert_data/const/artificial_tokens.py +0 -0
  24. cehrbert_data-0.0.2/src/cehrbert_data/decorators/__pycache__/__init__.cpython-311.pyc +0 -0
  25. cehrbert_data-0.0.2/src/cehrbert_data/decorators/__pycache__/patient_event_decorator.cpython-311.pyc +0 -0
  26. cehrbert_data-0.0.2/src/cehrbert_data/decorators/patient_event_decorator.py +0 -759
  27. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/.github/workflows/python-build.yml +0 -0
  28. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/.github/workflows/tests.yml +0 -0
  29. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/.gitignore +0 -0
  30. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/.pre-commit-config.yaml +0 -0
  31. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/LICENSE +0 -0
  32. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/README.md +0 -0
  33. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/pyproject.toml +0 -0
  34. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept/._SUCCESS.crc +0 -0
  35. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept/.part-00000-4b12270c-f6c8-4b59-8e0f-fd588bd79386-c000.snappy.parquet.crc +0 -0
  36. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept/.part-00003-4b12270c-f6c8-4b59-8e0f-fd588bd79386-c000.snappy.parquet.crc +0 -0
  37. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept/.part-00010-4b12270c-f6c8-4b59-8e0f-fd588bd79386-c000.snappy.parquet.crc +0 -0
  38. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept/_SUCCESS +0 -0
  39. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept/part-00000-4b12270c-f6c8-4b59-8e0f-fd588bd79386-c000.snappy.parquet +0 -0
  40. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept/part-00003-4b12270c-f6c8-4b59-8e0f-fd588bd79386-c000.snappy.parquet +0 -0
  41. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept/part-00010-4b12270c-f6c8-4b59-8e0f-fd588bd79386-c000.snappy.parquet +0 -0
  42. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_ancestor/._SUCCESS.crc +0 -0
  43. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_ancestor/.part-00000-eafbd8be-3337-46da-89d3-20f79c2565d4-c000.snappy.parquet.crc +0 -0
  44. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_ancestor/.part-00002-eafbd8be-3337-46da-89d3-20f79c2565d4-c000.snappy.parquet.crc +0 -0
  45. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_ancestor/.part-00006-eafbd8be-3337-46da-89d3-20f79c2565d4-c000.snappy.parquet.crc +0 -0
  46. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_ancestor/.part-00011-eafbd8be-3337-46da-89d3-20f79c2565d4-c000.snappy.parquet.crc +0 -0
  47. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_ancestor/.part-00013-eafbd8be-3337-46da-89d3-20f79c2565d4-c000.snappy.parquet.crc +0 -0
  48. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_ancestor/_SUCCESS +0 -0
  49. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_ancestor/part-00000-eafbd8be-3337-46da-89d3-20f79c2565d4-c000.snappy.parquet +0 -0
  50. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_ancestor/part-00002-eafbd8be-3337-46da-89d3-20f79c2565d4-c000.snappy.parquet +0 -0
  51. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_ancestor/part-00006-eafbd8be-3337-46da-89d3-20f79c2565d4-c000.snappy.parquet +0 -0
  52. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_ancestor/part-00011-eafbd8be-3337-46da-89d3-20f79c2565d4-c000.snappy.parquet +0 -0
  53. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_ancestor/part-00013-eafbd8be-3337-46da-89d3-20f79c2565d4-c000.snappy.parquet +0 -0
  54. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_relationship/._SUCCESS.crc +0 -0
  55. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_relationship/.part-00000-5752b472-8ba7-4189-ab69-8c92e46443aa-c000.snappy.parquet.crc +0 -0
  56. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_relationship/.part-00002-5752b472-8ba7-4189-ab69-8c92e46443aa-c000.snappy.parquet.crc +0 -0
  57. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_relationship/.part-00007-5752b472-8ba7-4189-ab69-8c92e46443aa-c000.snappy.parquet.crc +0 -0
  58. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_relationship/.part-00012-5752b472-8ba7-4189-ab69-8c92e46443aa-c000.snappy.parquet.crc +0 -0
  59. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_relationship/_SUCCESS +0 -0
  60. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_relationship/part-00000-5752b472-8ba7-4189-ab69-8c92e46443aa-c000.snappy.parquet +0 -0
  61. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_relationship/part-00002-5752b472-8ba7-4189-ab69-8c92e46443aa-c000.snappy.parquet +0 -0
  62. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_relationship/part-00007-5752b472-8ba7-4189-ab69-8c92e46443aa-c000.snappy.parquet +0 -0
  63. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/concept_relationship/part-00012-5752b472-8ba7-4189-ab69-8c92e46443aa-c000.snappy.parquet +0 -0
  64. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/condition_occurrence/._SUCCESS.crc +0 -0
  65. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/condition_occurrence/.part-00000-4eff03a1-cdcf-4c89-b0cd-9ce590b9b1eb-c000.snappy.parquet.crc +0 -0
  66. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/condition_occurrence/_SUCCESS +0 -0
  67. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/condition_occurrence/part-00000-4eff03a1-cdcf-4c89-b0cd-9ce590b9b1eb-c000.snappy.parquet +0 -0
  68. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/drug_exposure/._SUCCESS.crc +0 -0
  69. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/drug_exposure/.part-00000-10bbf1a4-a7da-416e-9703-58609c7edfad-c000.snappy.parquet.crc +0 -0
  70. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/drug_exposure/_SUCCESS +0 -0
  71. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/drug_exposure/part-00000-10bbf1a4-a7da-416e-9703-58609c7edfad-c000.snappy.parquet +0 -0
  72. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/observation_period/._SUCCESS.crc +0 -0
  73. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/observation_period/.part-00000-694316e5-cc95-49f1-9fad-5a7f377e2602-c000.snappy.parquet.crc +0 -0
  74. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/observation_period/_SUCCESS +0 -0
  75. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/observation_period/part-00000-694316e5-cc95-49f1-9fad-5a7f377e2602-c000.snappy.parquet +0 -0
  76. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/person/._SUCCESS.crc +0 -0
  77. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/person/.part-00000-7d789011-f361-48da-af6f-cfe102978b3a-c000.snappy.parquet.crc +0 -0
  78. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/person/_SUCCESS +0 -0
  79. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/person/part-00000-7d789011-f361-48da-af6f-cfe102978b3a-c000.snappy.parquet +0 -0
  80. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/procedure_occurrence/._SUCCESS.crc +0 -0
  81. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/procedure_occurrence/.part-00000-e73003c1-aed5-41c0-b2d4-eaccaccf044a-c000.snappy.parquet.crc +0 -0
  82. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/procedure_occurrence/_SUCCESS +0 -0
  83. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/procedure_occurrence/part-00000-e73003c1-aed5-41c0-b2d4-eaccaccf044a-c000.snappy.parquet +0 -0
  84. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/visit_occurrence/._SUCCESS.crc +0 -0
  85. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/visit_occurrence/.part-00000-e874b5f1-bf9e-4cb9-93bf-c309a47b0476-c000.snappy.parquet.crc +0 -0
  86. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/visit_occurrence/_SUCCESS +0 -0
  87. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/sample_data/omop_sample/visit_occurrence/part-00000-e874b5f1-bf9e-4cb9-93bf-c309a47b0476-c000.snappy.parquet +0 -0
  88. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/setup.cfg +0 -0
  89. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/__init__.py +0 -0
  90. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/__init__.py +0 -0
  91. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/apps/__init__.py +0 -0
  92. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/apps/generate_concept_similarity_table.py +0 -0
  93. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/apps/generate_hierarchical_bert_training_data.py +0 -0
  94. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/apps/generate_included_concept_list.py +0 -0
  95. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/apps/generate_information_content.py +0 -0
  96. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/cohorts/__init__.py +0 -0
  97. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/cohorts/atrial_fibrillation.py +0 -0
  98. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/cohorts/cabg.py +0 -0
  99. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/cohorts/coronary_artery_disease.py +0 -0
  100. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/cohorts/covid.py +0 -0
  101. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/cohorts/covid_inpatient.py +0 -0
  102. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/cohorts/death.py +0 -0
  103. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/cohorts/heart_failure.py +0 -0
  104. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/cohorts/ischemic_stroke.py +0 -0
  105. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/cohorts/last_visit_discharged_home.py +0 -0
  106. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/cohorts/query_builder.py +0 -0
  107. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/cohorts/type_two_diabietes.py +0 -0
  108. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/cohorts/ventilation.py +0 -0
  109. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/config/__init__.py +0 -0
  110. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/config/output_names.py +0 -0
  111. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/const/__init__.py +0 -0
  112. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/const/__pycache__/__init__.cpython-311.pyc +0 -0
  113. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/const/__pycache__/common.cpython-311.pyc +0 -0
  114. {cehrbert_data-0.0.2/src/cehrbert_data/decorators → cehrbert_data-0.0.4/src/cehrbert_data/prediction_cohorts}/__init__.py +0 -0
  115. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/prediction_cohorts/afib_ischemic_stroke.py +0 -0
  116. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/prediction_cohorts/cad_cabg_cohort.py +0 -0
  117. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/prediction_cohorts/cad_hf_cohort.py +0 -0
  118. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/prediction_cohorts/copd_readmission.py +0 -0
  119. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/prediction_cohorts/covid_death.py +0 -0
  120. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/prediction_cohorts/covid_ventilation.py +0 -0
  121. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/prediction_cohorts/discharge_home_death.py +0 -0
  122. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/prediction_cohorts/hf_readmission.py +0 -0
  123. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/prediction_cohorts/hospitalization.py +0 -0
  124. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/prediction_cohorts/hospitalization_mortality.py +0 -0
  125. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/prediction_cohorts/t2dm_hf_cohort.py +0 -0
  126. {cehrbert_data-0.0.2/src/cehrbert_data/prediction_cohorts → cehrbert_data-0.0.4/src/cehrbert_data/queries}/__init__.py +0 -0
  127. {cehrbert_data-0.0.2/src/cehrbert_data/queries → cehrbert_data-0.0.4/src/cehrbert_data/tools}/__init__.py +0 -0
  128. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data/tools/download_omop_tables.py +0 -0
  129. {cehrbert_data-0.0.2/src/cehrbert_data/tools → cehrbert_data-0.0.4/src/cehrbert_data/utils}/__init__.py +0 -0
  130. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data.egg-info/dependency_links.txt +0 -0
  131. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data.egg-info/requires.txt +0 -0
  132. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/src/cehrbert_data.egg-info/top_level.txt +0 -0
  133. {cehrbert_data-0.0.2/src/cehrbert_data/utils → cehrbert_data-0.0.4/tests}/__init__.py +0 -0
  134. {cehrbert_data-0.0.2/tests → cehrbert_data-0.0.4/tests/integration_tests}/__init__.py +0 -0
  135. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/tests/integration_tests/test_hf_readmission.py +0 -0
  136. {cehrbert_data-0.0.2 → cehrbert_data-0.0.4}/tests/pyspark_test_base.py +0 -0
  137. {cehrbert_data-0.0.2/tests/integration_tests → cehrbert_data-0.0.4/tests/unit_tests}/__init__.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: cehrbert_data
3
- Version: 0.0.2
3
+ Version: 0.0.4
4
4
  Summary: The Spark ETL tools for generating the CEHR-BERT and CEHR-GPT pre-training and finetuning data
5
5
  Author-email: Chao Pang <chaopang229@gmail.com>, Xinzhuo Jiang <xj2193@cumc.columbia.edu>, Krishna Kalluri <kk3326@cumc.columbia.edu>, Nishanth Parameshwar Pavinkurve <np2689@cumc.columbia.edu>, Karthik Natarajan <kn2174@cumc.columbia.edu>
6
6
  License: MIT License
@@ -0,0 +1,140 @@
1
+ import argparse
2
+ import os
3
+
4
+ from pyspark.sql import SparkSession, DataFrame
5
+ from pyspark.sql import functions as F
6
+ from pyspark.sql.window import Window
7
+
8
+ from cehrbert_data.const.common import (
9
+ CONCEPT,
10
+ MEASUREMENT,
11
+ REQUIRED_MEASUREMENT,
12
+ NUMERIC_MEASUREMENT_STATS
13
+ )
14
+ from cehrbert_data.utils.spark_utils import preprocess_domain_table
15
+ from cehrbert_data.queries.measurement_queries import (
16
+ LAB_PREVALENCE_QUERY,
17
+ MEASUREMENT_UNIT_STATS_QUERY
18
+ )
19
+
20
+
21
+ def main(
22
+ input_folder,
23
+ output_folder,
24
+ num_of_numeric_labs,
25
+ num_of_categorical_labs,
26
+ min_num_of_patients
27
+ ):
28
+ spark = SparkSession.builder.appName("Generate required labs").getOrCreate()
29
+
30
+ # Load measurement as a dataframe in pyspark
31
+ measurement = preprocess_domain_table(spark, input_folder, MEASUREMENT)
32
+ # Load concept as a dataframe in pyspark
33
+ concept = preprocess_domain_table(spark, input_folder, CONCEPT)
34
+ # Create the local measurement view
35
+ measurement.createOrReplaceTempView(MEASUREMENT)
36
+ # Create the local concept view
37
+ concept.createOrReplaceTempView(CONCEPT)
38
+ # Create the
39
+ required_lab_dataframe = generate_required_labs(
40
+ spark, num_of_numeric_labs, num_of_categorical_labs, min_num_of_patients
41
+ )
42
+ required_lab_dataframe.write.mode("overwrite").parquet(
43
+ os.path.join(output_folder, REQUIRED_MEASUREMENT)
44
+ )
45
+ # Reload the dataframe from the disk
46
+ required_lab_dataframe = spark.read.parquet(
47
+ os.path.join(output_folder, REQUIRED_MEASUREMENT)
48
+ )
49
+ required_lab_dataframe.createOrReplaceTempView(REQUIRED_MEASUREMENT)
50
+ numeric_measurement_stats_dataframe = spark.sql(MEASUREMENT_UNIT_STATS_QUERY)
51
+ numeric_measurement_stats_dataframe.write.mode("overwrite").parquet(
52
+ os.path.join(output_folder, NUMERIC_MEASUREMENT_STATS)
53
+ )
54
+
55
+
56
+ def generate_required_labs(
57
+ spark: SparkSession,
58
+ num_of_numeric_labs: int,
59
+ num_of_categorical_labs: int,
60
+ min_num_of_patients: int
61
+ ) -> DataFrame:
62
+
63
+ prevalent_labs = spark.sql(LAB_PREVALENCE_QUERY)
64
+ prevalent_labs = prevalent_labs.where(F.col("person_count") >= min_num_of_patients)
65
+ # Cache the dataframe for faster computation in the below transformations
66
+ prevalent_labs.cache()
67
+ prevalent_numeric_labs = (
68
+ prevalent_labs.withColumn("is_numeric", F.col("numeric_percentage") >= 0.5)
69
+ .where("is_numeric")
70
+ .withColumn("rn", F.row_number().over(Window.orderBy(F.desc("freq"))))
71
+ .where(F.col("rn") <= num_of_numeric_labs)
72
+ .drop("rn")
73
+ )
74
+ prevalent_categorical_labs = (
75
+ prevalent_labs.withColumn("is_categorical", F.col("categorical_percentage") >= 0.5)
76
+ .where("is_categorical")
77
+ .withColumn("is_numeric", ~F.col("is_categorical"))
78
+ .withColumn("rn", F.row_number().over(Window.orderBy(F.desc("freq"))))
79
+ .where(F.col("rn") <= num_of_categorical_labs)
80
+ .drop("is_categorical")
81
+ .drop("rn")
82
+ )
83
+ return prevalent_numeric_labs.unionAll(prevalent_categorical_labs)
84
+
85
+
86
+ if __name__ == "__main__":
87
+ parser = argparse.ArgumentParser(description="Arguments for generate " "required labs to be included")
88
+ parser.add_argument(
89
+ "-i",
90
+ "--input_folder",
91
+ dest="input_folder",
92
+ action="store",
93
+ help="The path for your input_folder where the raw data is",
94
+ required=True,
95
+ )
96
+ parser.add_argument(
97
+ "-o",
98
+ "--output_folder",
99
+ dest="output_folder",
100
+ action="store",
101
+ help="The path for your output_folder",
102
+ required=True,
103
+ )
104
+ parser.add_argument(
105
+ "--num_of_numeric_labs",
106
+ dest="num_of_numeric_labs",
107
+ action="store",
108
+ type=int,
109
+ default=100,
110
+ help="The top most prevalent numeric labs to be included",
111
+ required=False,
112
+ )
113
+ parser.add_argument(
114
+ "--num_of_categorical_labs",
115
+ dest="num_of_categorical_labs",
116
+ action="store",
117
+ type=int,
118
+ default=100,
119
+ help="The top most prevalent categorical labs to be included",
120
+ required=False,
121
+ )
122
+ parser.add_argument(
123
+ "--min_num_of_patients",
124
+ dest="min_num_of_patients",
125
+ action="store",
126
+ type=int,
127
+ default=0,
128
+ help="Min no.of patients linked to concepts to be included",
129
+ required=False,
130
+ )
131
+
132
+ ARGS = parser.parse_args()
133
+
134
+ main(
135
+ ARGS.input_folder,
136
+ ARGS.output_folder,
137
+ ARGS.num_of_numeric_labs,
138
+ ARGS.num_of_categorical_labs,
139
+ ARGS.min_num_of_patients
140
+ )
@@ -8,39 +8,46 @@ from pyspark.sql import SparkSession
8
8
  from pyspark.sql import functions as F
9
9
  from pyspark.sql.window import Window
10
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
11
+ from cehrbert_data.const.common import (
12
+ PERSON,
13
+ VISIT_OCCURRENCE,
14
+ DEATH,
15
+ MEASUREMENT
16
+ )
17
+ from cehrbert_data.decorators import AttType
13
18
  from cehrbert_data.utils.spark_utils import (
14
19
  create_sequence_data,
15
20
  create_sequence_data_with_att,
16
21
  join_domain_tables,
17
22
  preprocess_domain_table,
18
- process_measurement,
23
+ get_measurement_table,
19
24
  validate_table_names,
20
25
  )
26
+ from cehrbert_data.utils.logging_utils import add_console_logging
21
27
 
22
28
 
23
29
  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,
30
+ input_folder,
31
+ output_folder,
32
+ domain_table_list,
33
+ date_filter,
34
+ include_visit_type,
35
+ is_new_patient_representation,
36
+ exclude_visit_tokens,
37
+ is_classic_bert,
38
+ include_prolonged_stay,
39
+ include_concept_list: bool,
40
+ gpt_patient_sequence: bool,
41
+ apply_age_filter: bool,
42
+ include_death: bool,
43
+ att_type: AttType,
44
+ include_sequence_information_content: bool = False,
45
+ exclude_demographic: bool = False,
46
+ use_age_group: bool = False,
47
+ with_drug_rollup: bool = True,
48
+ include_inpatient_hour_token: bool = False,
49
+ continue_from_events: bool = False,
50
+ refresh_measurement: bool = False
44
51
  ):
45
52
  spark = SparkSession.builder.appName("Generate CEHR-BERT Training Data").getOrCreate()
46
53
 
@@ -63,6 +70,7 @@ def main(
63
70
  f"exclude_demographic: {exclude_demographic}\n"
64
71
  f"use_age_group: {use_age_group}\n"
65
72
  f"with_drug_rollup: {with_drug_rollup}\n"
73
+ f"refresh_measurement: {refresh_measurement}\n"
66
74
  )
67
75
 
68
76
  domain_tables = []
@@ -117,17 +125,16 @@ def main(
117
125
 
118
126
  # Process the measurement table if exists
119
127
  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
-
128
+ processed_measurement = get_measurement_table(
129
+ spark,
130
+ input_folder,
131
+ refresh=refresh_measurement
132
+ )
126
133
  if patient_events:
127
134
  # Union all measurement records together with other domain records
128
- patient_events = patient_events.unionByName(scaled_measurement)
135
+ patient_events = patient_events.unionByName(processed_measurement)
129
136
  else:
130
- patient_events = scaled_measurement
137
+ patient_events = processed_measurement
131
138
 
132
139
  patient_events = (
133
140
  patient_events.join(visit_occurrence_person, "visit_occurrence_id")
@@ -206,17 +213,21 @@ def main(
206
213
  patient_splits_folder = os.path.join(input_folder, "patient_splits")
207
214
  if os.path.exists(patient_splits_folder):
208
215
  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")
216
+ temp_folder = os.path.join(output_folder, "patient_sequence", "temp")
217
+ sequence_data.join(
218
+ patient_splits.select("person_id", "split"),
219
+ "person_id"
220
+ ).write.mode("overwrite").parquet(
221
+ temp_folder
211
222
  )
212
- sequence_data = spark.read.parquet(os.path.join(output_folder, "patient_sequence", "temp"))
223
+ sequence_data = spark.read.parquet(temp_folder)
213
224
  sequence_data.where('split="train"').write.mode("overwrite").parquet(
214
225
  os.path.join(output_folder, "patient_sequence/train")
215
226
  )
216
227
  sequence_data.where('split="test"').write.mode("overwrite").parquet(
217
228
  os.path.join(output_folder, "patient_sequence/test")
218
229
  )
219
- shutil.rmtree(os.path.join(output_folder, "patient_sequence", "temp"))
230
+ shutil.rmtree(temp_folder)
220
231
  else:
221
232
  sequence_data.write.mode("overwrite").parquet(os.path.join(output_folder, "patient_sequence"))
222
233
 
@@ -304,7 +315,16 @@ if __name__ == "__main__":
304
315
  dest="include_inpatient_hour_token",
305
316
  action="store_true",
306
317
  )
307
- parser.add_argument("--continue_from_events", dest="continue_from_events", action="store_true")
318
+ parser.add_argument(
319
+ "--continue_from_events",
320
+ dest="continue_from_events",
321
+ action="store_true"
322
+ )
323
+ parser.add_argument(
324
+ "--refresh_measurement",
325
+ dest="refresh_measurement",
326
+ action="store_true"
327
+ )
308
328
  parser.add_argument(
309
329
  "--att_type",
310
330
  dest="att_type",
@@ -314,6 +334,9 @@ if __name__ == "__main__":
314
334
 
315
335
  ARGS = parser.parse_args()
316
336
 
337
+ # Enable logging
338
+ add_console_logging()
339
+
317
340
  main(
318
341
  ARGS.input_folder,
319
342
  ARGS.output_folder,
@@ -334,4 +357,5 @@ if __name__ == "__main__":
334
357
  with_drug_rollup=ARGS.with_drug_rollup,
335
358
  include_inpatient_hour_token=ARGS.include_inpatient_hour_token,
336
359
  continue_from_events=ARGS.continue_from_events,
360
+ refresh_measurement=ARGS.refresh_measurement,
337
361
  )
@@ -10,9 +10,9 @@ from pyspark.sql import DataFrame, SparkSession
10
10
  from pyspark.sql import functions as F
11
11
  from pyspark.sql.window import Window
12
12
 
13
- from cehrbert_data.decorators.patient_event_decorator import AttType
13
+ from cehrbert_data.decorators import AttType
14
+ from cehrbert_data.const.common import VISIT_OCCURRENCE
14
15
  from cehrbert_data.utils.spark_utils import (
15
- VISIT_OCCURRENCE,
16
16
  build_ancestry_table_for,
17
17
  create_concept_frequency_data,
18
18
  create_hierarchical_sequence_data,
@@ -22,6 +22,7 @@ from cehrbert_data.utils.spark_utils import (
22
22
  get_descendant_concept_ids,
23
23
  preprocess_domain_table,
24
24
  )
25
+ from cehrbert_data.utils.logging_utils import add_console_logging
25
26
 
26
27
  from ..cohorts.query_builder import ENTRY_COHORT, NEGATIVE_COHORT, QueryBuilder
27
28
 
@@ -97,6 +98,7 @@ class BaseCohortBuilder(ABC):
97
98
  age_upper_bound: int,
98
99
  prior_observation_period: int,
99
100
  post_observation_period: int,
101
+ continue_job: bool = False
100
102
  ):
101
103
 
102
104
  self._query_builder = query_builder
@@ -110,6 +112,7 @@ class BaseCohortBuilder(ABC):
110
112
  self._post_observation_period = post_observation_period
111
113
  cohort_name = re.sub("[^a-z0-9]+", "_", self._query_builder.get_cohort_name().lower())
112
114
  self._output_data_folder = os.path.join(self._output_folder, cohort_name)
115
+ self._continue_job = continue_job
113
116
 
114
117
  self.get_logger().info(
115
118
  f"query_builder: {query_builder}\n"
@@ -121,6 +124,7 @@ class BaseCohortBuilder(ABC):
121
124
  f"age_upper_bound: {age_upper_bound}\n"
122
125
  f"prior_observation_period: {prior_observation_period}\n"
123
126
  f"post_observation_period: {post_observation_period}\n"
127
+ f"continue_job: {continue_job}\n"
124
128
  )
125
129
 
126
130
  # Validate the age range, observation_window and prediction_window
@@ -187,6 +191,11 @@ class BaseCohortBuilder(ABC):
187
191
 
188
192
  def build(self):
189
193
  """Build the cohort and write the dataframe as parquet files to _output_data_folder."""
194
+
195
+ # Check whether the cohort has been generated
196
+ if self._continue_job and self.cohort_exists():
197
+ return self
198
+
190
199
  cohort = self.create_cohort()
191
200
 
192
201
  cohort = self._apply_observation_period(cohort)
@@ -201,6 +210,13 @@ class BaseCohortBuilder(ABC):
201
210
 
202
211
  return self
203
212
 
213
+ def cohort_exists(self) -> bool:
214
+ try:
215
+ self.load_cohort()
216
+ return True
217
+ except Exception:
218
+ return False
219
+
204
220
  def load_cohort(self):
205
221
  return self.spark.read.parquet(self._output_data_folder)
206
222
 
@@ -276,7 +292,9 @@ class NestedCohortBuilder:
276
292
  exclude_visit_tokens: bool = False,
277
293
  is_feature_concept_frequency: bool = False,
278
294
  is_roll_up_concept: bool = False,
295
+ is_drug_roll_up_concept: bool = True,
279
296
  include_concept_list: bool = True,
297
+ refresh_measurement: bool = False,
280
298
  is_new_patient_representation: bool = False,
281
299
  gpt_patient_sequence: bool = False,
282
300
  is_hierarchical_bert: bool = False,
@@ -312,6 +330,7 @@ class NestedCohortBuilder:
312
330
  self._classic_bert_seq = classic_bert_seq
313
331
  self._is_feature_concept_frequency = is_feature_concept_frequency
314
332
  self._is_roll_up_concept = is_roll_up_concept
333
+ self._is_drug_roll_up_concept = is_drug_roll_up_concept
315
334
  self._is_new_patient_representation = is_new_patient_representation
316
335
  self._gpt_patient_sequence = gpt_patient_sequence
317
336
  self._is_hierarchical_bert = is_hierarchical_bert
@@ -320,6 +339,7 @@ class NestedCohortBuilder:
320
339
  self._is_questionable_outcome_existed = is_questionable_outcome_existed
321
340
  self._is_prediction_window_unbounded = is_prediction_window_unbounded
322
341
  self._include_concept_list = include_concept_list
342
+ self._refresh_measurement = refresh_measurement
323
343
  self._allow_measurement_only = allow_measurement_only
324
344
  self._output_data_folder = os.path.join(
325
345
  self._output_folder, re.sub("[^a-z0-9]+", "_", self._cohort_name.lower())
@@ -347,6 +367,7 @@ class NestedCohortBuilder:
347
367
  f"allow_measurement_only: {allow_measurement_only}\n"
348
368
  f"is_feature_concept_frequency: {is_feature_concept_frequency}\n"
349
369
  f"is_roll_up_concept: {is_roll_up_concept}\n"
370
+ f"is_drug_roll_up_concept: {is_drug_roll_up_concept}\n"
350
371
  f"is_new_patient_representation: {is_new_patient_representation}\n"
351
372
  f"gpt_patient_sequence: {gpt_patient_sequence}\n"
352
373
  f"is_hierarchical_bert: {is_hierarchical_bert}\n"
@@ -355,6 +376,7 @@ class NestedCohortBuilder:
355
376
  f"is_remove_index_prediction_starts: {is_remove_index_prediction_starts}\n"
356
377
  f"is_prediction_window_unbounded: {is_prediction_window_unbounded}\n"
357
378
  f"include_concept_list: {include_concept_list}\n"
379
+ f"refresh_measurement: {refresh_measurement}\n"
358
380
  f"is_observation_window_unbounded: {is_observation_window_unbounded}\n"
359
381
  f"is_population_estimation: {is_population_estimation}\n"
360
382
  f"att_type: {att_type}\n"
@@ -494,11 +516,21 @@ class NestedCohortBuilder:
494
516
  .where(F.col("num_of_concepts") >= self._num_of_concepts)
495
517
  )
496
518
 
497
- # Add time_to_event
498
- cohort = cohort.withColumn(
499
- "study_end_date",
500
- F.coalesce(F.col("outcome_date"), F.date_add(cohort.index_date, self._prediction_window))
501
- )
519
+ if self._is_prediction_window_unbounded:
520
+ observation_period = self._dependency_dict[OBSERVATION_PERIOD]
521
+ # Add time_to_event
522
+ cohort = cohort.join(
523
+ observation_period.select("person_id", "observation_period_end_date"), "person_id"
524
+ ).withColumn(
525
+ "study_end_date",
526
+ F.coalesce(F.col("outcome_date"), F.col("observation_period_end_date"))
527
+ ).drop("observation_period_end_date")
528
+ else:
529
+ # Add time_to_event
530
+ cohort = cohort.withColumn(
531
+ "study_end_date",
532
+ F.coalesce(F.col("outcome_date"), F.date_add(cohort.index_date, self._prediction_window))
533
+ )
502
534
  cohort = cohort.withColumn("time_to_event", F.datediff("study_end_date", "index_date"))
503
535
 
504
536
  # if patient_splits is provided, we will
@@ -526,12 +558,14 @@ class NestedCohortBuilder:
526
558
  """
527
559
  # Extract all ehr records for the patients
528
560
  ehr_records = extract_ehr_records(
529
- self.spark,
530
- self._input_folder,
531
- self._ehr_table_list,
532
- self._include_visit_type,
533
- self._is_roll_up_concept,
534
- self._include_concept_list,
561
+ spark=self.spark,
562
+ input_folder=self._input_folder,
563
+ domain_table_list=self._ehr_table_list,
564
+ include_visit_type=self._include_visit_type,
565
+ with_diagnosis_rollup=self._is_roll_up_concept,
566
+ with_drug_rollup=self._is_drug_roll_up_concept,
567
+ include_concept_list=self._include_concept_list,
568
+ refresh_measurement=self._refresh_measurement,
535
569
  )
536
570
 
537
571
  # Duplicate the records for cohorts that allow multiple entries
@@ -644,54 +678,31 @@ def create_prediction_cohort(
644
678
  :param ehr_table_list:
645
679
  :return:
646
680
  """
647
- cohort_name = spark_args.cohort_name
648
- input_folder = spark_args.input_folder
649
- output_folder = spark_args.output_folder
650
- date_lower_bound = spark_args.date_lower_bound
651
- date_upper_bound = spark_args.date_upper_bound
652
- age_lower_bound = spark_args.age_lower_bound
653
- age_upper_bound = spark_args.age_upper_bound
654
- observation_window = spark_args.observation_window
655
- prediction_start_days = spark_args.prediction_start_days
656
- prediction_window = spark_args.prediction_window
657
- hold_off_window = spark_args.hold_off_window
658
- num_of_visits = spark_args.num_of_visits
659
- num_of_concepts = spark_args.num_of_concepts
660
- include_visit_type = spark_args.include_visit_type
661
- exclude_visit_tokens = spark_args.exclude_visit_tokens
662
- is_feature_concept_frequency = spark_args.is_feature_concept_frequency
663
- is_roll_up_concept = spark_args.is_roll_up_concept
664
- is_window_post_index = spark_args.is_window_post_index
665
- is_new_patient_representation = spark_args.is_new_patient_representation
666
- is_hierarchical_bert = spark_args.is_hierarchical_bert
667
- classic_bert_seq = spark_args.classic_bert_seq
668
- is_first_time_outcome = spark_args.is_first_time_outcome
669
- is_prediction_window_unbounded = spark_args.is_prediction_window_unbounded
670
- is_observation_window_unbounded = spark_args.is_observation_window_unbounded
671
- # If the outcome negative query exists, that means we need to remove those questionable
672
- # outcomes from the target cohort
673
- is_questionable_outcome_existed = outcome_query_builder.get_negative_query() is not None
674
-
675
- # Do we want to remove those records whose outcome occur between index_date and the start of
676
- # the prediction window
677
- is_remove_index_prediction_starts = spark_args.is_remove_index_prediction_starts
681
+ # Add logging to spark application output when enable_logging is set to True
682
+ if spark_args.enable_logging:
683
+ add_console_logging()
678
684
 
679
685
  # Toggle the prior/post observation_period depending on the is_window_post_index flag
680
- prior_observation_period = 0 if is_window_post_index else observation_window + hold_off_window
681
- post_observation_period = observation_window + hold_off_window if is_window_post_index else 0
686
+ prior_observation_period = (
687
+ 0 if spark_args.is_window_post_index else spark_args.observation_window + spark_args.hold_off_window
688
+ )
689
+ post_observation_period = (
690
+ spark_args.observation_window + spark_args.hold_off_window if spark_args.is_window_post_index else 0
691
+ )
682
692
 
683
693
  # Generate the target cohort
684
694
  target_cohort = (
685
695
  BaseCohortBuilder(
686
696
  query_builder=target_query_builder,
687
- input_folder=input_folder,
688
- output_folder=output_folder,
689
- date_lower_bound=date_lower_bound,
690
- date_upper_bound=date_upper_bound,
691
- age_lower_bound=age_lower_bound,
692
- age_upper_bound=age_upper_bound,
697
+ input_folder=spark_args.input_folder,
698
+ output_folder=spark_args.output_folder,
699
+ date_lower_bound=spark_args.date_lower_bound,
700
+ date_upper_bound=spark_args.date_upper_bound,
701
+ age_lower_bound=spark_args.age_lower_bound,
702
+ age_upper_bound=spark_args.age_upper_bound,
693
703
  prior_observation_period=prior_observation_period,
694
704
  post_observation_period=post_observation_period,
705
+ continue_job=spark_args.continue_job
695
706
  )
696
707
  .build()
697
708
  .load_cohort()
@@ -701,12 +712,12 @@ def create_prediction_cohort(
701
712
  outcome_cohort = (
702
713
  BaseCohortBuilder(
703
714
  query_builder=outcome_query_builder,
704
- input_folder=input_folder,
705
- output_folder=output_folder,
706
- date_lower_bound=date_lower_bound,
707
- date_upper_bound=date_upper_bound,
708
- age_lower_bound=age_lower_bound,
709
- age_upper_bound=age_upper_bound,
715
+ input_folder=spark_args.input_folder,
716
+ output_folder=spark_args.output_folder,
717
+ date_lower_bound=spark_args.date_lower_bound,
718
+ date_upper_bound=spark_args.date_upper_bound,
719
+ age_lower_bound=spark_args.age_lower_bound,
720
+ age_upper_bound=spark_args.age_upper_bound,
710
721
  prior_observation_period=0,
711
722
  post_observation_period=0,
712
723
  )
@@ -715,35 +726,41 @@ def create_prediction_cohort(
715
726
  )
716
727
 
717
728
  NestedCohortBuilder(
718
- cohort_name=cohort_name,
719
- input_folder=input_folder,
720
- output_folder=output_folder,
729
+ cohort_name=spark_args.cohort_name,
730
+ input_folder=spark_args.input_folder,
731
+ output_folder=spark_args.output_folder,
721
732
  patient_splits_folder=spark_args.patient_splits_folder,
722
733
  target_cohort=target_cohort,
723
734
  outcome_cohort=outcome_cohort,
724
735
  ehr_table_list=ehr_table_list,
725
- observation_window=observation_window,
726
- hold_off_window=hold_off_window,
727
- prediction_start_days=prediction_start_days,
728
- prediction_window=prediction_window,
729
- num_of_visits=num_of_visits,
730
- num_of_concepts=num_of_concepts,
731
- is_window_post_index=is_window_post_index,
732
- include_visit_type=include_visit_type,
733
- exclude_visit_tokens=exclude_visit_tokens,
736
+ observation_window=spark_args.observation_window,
737
+ hold_off_window=spark_args.hold_off_window,
738
+ prediction_start_days=spark_args.prediction_start_days,
739
+ prediction_window=spark_args.prediction_window,
740
+ num_of_visits=spark_args.num_of_visits,
741
+ num_of_concepts=spark_args.num_of_concepts,
742
+ is_window_post_index=spark_args.is_window_post_index,
743
+ include_visit_type=spark_args.include_visit_type,
744
+ exclude_visit_tokens=spark_args.exclude_visit_tokens,
734
745
  allow_measurement_only=spark_args.allow_measurement_only,
735
- is_feature_concept_frequency=is_feature_concept_frequency,
736
- is_roll_up_concept=is_roll_up_concept,
746
+ is_feature_concept_frequency=spark_args.is_feature_concept_frequency,
747
+ is_roll_up_concept=spark_args.is_roll_up_concept,
748
+ is_drug_roll_up_concept=spark_args.is_drug_roll_up_concept,
737
749
  include_concept_list=spark_args.include_concept_list,
738
- is_new_patient_representation=is_new_patient_representation,
750
+ refresh_measurement=spark_args.refresh_measurement,
751
+ is_new_patient_representation=spark_args.is_new_patient_representation,
739
752
  gpt_patient_sequence=spark_args.gpt_patient_sequence,
740
- is_hierarchical_bert=is_hierarchical_bert,
741
- classic_bert_seq=classic_bert_seq,
742
- is_first_time_outcome=is_first_time_outcome,
743
- is_questionable_outcome_existed=is_questionable_outcome_existed,
744
- is_prediction_window_unbounded=is_prediction_window_unbounded,
745
- is_remove_index_prediction_starts=is_remove_index_prediction_starts,
746
- is_observation_window_unbounded=is_observation_window_unbounded,
753
+ is_hierarchical_bert=spark_args.is_hierarchical_bert,
754
+ classic_bert_seq=spark_args.classic_bert_seq,
755
+ is_first_time_outcome=spark_args.is_first_time_outcome,
756
+ # If the outcome negative query exists, that means we need to remove those questionable
757
+ # outcomes from the target cohort
758
+ is_questionable_outcome_existed=outcome_query_builder.get_negative_query() is not None,
759
+ is_prediction_window_unbounded=spark_args.is_prediction_window_unbounded,
760
+ # Do we want to remove those records whose outcome occur between index_date
761
+ # and the start of the prediction window
762
+ is_remove_index_prediction_starts=spark_args.is_remove_index_prediction_starts,
763
+ is_observation_window_unbounded=spark_args.is_observation_window_unbounded,
747
764
  is_population_estimation=spark_args.is_population_estimation,
748
765
  att_type=AttType(spark_args.att_type),
749
766
  exclude_demographic=spark_args.exclude_demographic,
@@ -0,0 +1,3 @@
1
+ VS_TOKEN = "[VS]"
2
+ VE_TOKEN = "[VE]"
3
+ DEATH_TOKEN = "[DEATH]"
@@ -6,6 +6,7 @@ DRUG_EXPOSURE = "drug_exposure"
6
6
  DEVICE_EXPOSURE = "device_exposure"
7
7
  OBSERVATION = "observation"
8
8
  MEASUREMENT = "measurement"
9
+ PROCESSED_MEASUREMENT = "processed_measurement"
9
10
  CATEGORICAL_MEASUREMENT = "categorical_measurement"
10
11
  OBSERVATION_PERIOD = "observation_period"
11
12
  DEATH = "death"
@@ -23,6 +24,10 @@ CDM_TABLES = [
23
24
  DEATH,
24
25
  ]
25
26
  REQUIRED_MEASUREMENT = "required_measurement"
27
+ NUMERIC_MEASUREMENT_STATS = "numeric_measurement_stats"
26
28
  UNKNOWN_CONCEPT = "[UNKNOWN]"
29
+ NA = "N/A"
27
30
  CONCEPT = "concept"
28
31
  CONCEPT_ANCESTOR = "concept_ancestor"
32
+ MEASUREMENT_QUESTION_PREFIX = "1-Question:"
33
+ MEASUREMENT_ANSWER_PREFIX = "2-Answer:"
@@ -0,0 +1,5 @@
1
+ from .artificial_time_token_decorator import AttEventDecorator
2
+ from .death_event_decorator import DeathEventDecorator
3
+ from .clinical_event_decorator import ClinicalEventDecorator
4
+ from .demographic_event_decorator import DemographicEventDecorator
5
+ from .patient_event_decorator_base import time_token_func, get_att_function, AttType