Collaborative (N3C)
Sijia Liu1,*, Andrew Wen1,*, Liwei Wang1,*, Huan He1,*, Sunyang Fu1,*, Robert Miller2, Andrew Williams2, Daniel Harris3, Ramakanth Kavuluru3, Mei Liu4, Noor Abu-el-rub4, Dalton Schutte5, Rui Zhang5, Masoud Rouhizadeh6, John D. Osborne7, Yongqun He8, Umit Topaloglu9, Stephanie S Hong10, Joel H Saltz11, Thomas Schaffter12, Emily Pfaff13, Christopher G. Chute10, Tim Duong14, Melissa A. Haendel15, Rafael Fuentes16, Peter Szolovits17, Hua Xu18, Hongfang Liu1, National COVID Cohort Collaborative (N3C) Natural Language Processing (NLP) Subgroup, National
Covid Cohort Collaborative (N3C)
2. Tufts Clinical and Translational Science Institute, Tufts Medical Center. Cities. 9. Wake Forest Baptist Medical Center.
12. Sage Bionetwork. 14. Albert Einstein College of Medicine. 16. Alex Informatics.
Technology.
Abstract
Despite recent methodology advancements in clinical natural language processing (NLP), adoption of clinical NLP models within the clinical and translational research community remains hindered by issues with ETL process heterogeneity and human factor variations. In this study, we proposed an open NLP development framework with the aim of addressing these issues. The viability of such a platform was evaluated on a COVID-19 use case through sites participating in the National COVID Cohort Collaborative (N3C). As part of our assessment of the impact of single vs. multi-site NLP algorithm development, we evaluated the performance of both an NLP ruleset developed solely using a single site’s clinical narratives as well as one further refined using a synthetic derived dataset sourced from three sites (Mayo, UKen, and UMN). The single-site ruleset resulted in performances of 0.876, 0.706, and 0.694 in F-scores for Mayo, Minnesota, and Kentucky test datasets, respectively, while the multi-site NLP ruleset improved performances to 0.884, 0.769 and 0.806. The results of our use case test run inform us of the importance of a multi- site federated development, evaluation, and implementation framework. As such, we aim to meet this need with our framework by providing the tools necessary to conduct NLP development in a collaborative manner through consensus building, process coordination, and best practice sharing.
Introduction
Over the past decade, Electronic Health Record (EHR) systems have been increasingly implemented at US healthcare institutions. Large amounts of detailed longitudinal patient information, including lab tests, medications, disease status, and treatment outcomes, have consequently been accumulated and made electronically available. These large clinical databases are valuable data sources for clinical and translational research. As a result, major initiatives have been established to exploit this crucial resource, including the Clinical and Translational Science Awards (CTSA) Program’s National Center for Data to Health (CD2H)/National COVID Cohort Collaborative (N3C)1, 2, the Electronic Medical Records and Genomics (eMERGE) Network3, the Patient-Centered Outcomes Research Institute’s (PCORI) Clinical Research Networks (CRNs)4, the NIH All of Us Research Program5, and the Observational Health Data Science and Informatics (OHDSI) Consortia with demonstrated successes6, 7, 8, 9.
One common challenge faced by those initiatives is, however, the prevalence of clinical information embedded in unstructured text10. Compared to structured data entry, text is a more conventional way in the healthcare environment to document impressions, clinical findings, assessments, and care plans. Even with the advent of sophisticated EHR systems, studies have shown that capturing health information fully in structured format through data entry is unlikely to happen and a blended model where physicians use templates when and where possible and dictate the details of a patient visit in text11.
Natural language processing (NLP) has been promoted as having a great potential to extract information from text12. NLP algorithms can generally be categorized into using either symbolic or statistical methods13. Since the turn of the century, machine learning algorithms (i.e., statistical NLP) have gained increased prominence in clinical NLP research14. Nevertheless, a substantial portion of clinical NLP use cases leverages symbolic techniques given that dictionary or rule-based methodologies suffice to meet the information needs of many clinical applications under specific use cases. In the context of EHR-based clinical research, NLP has been leveraged to assist information extraction and knowledge conversion at different stages of research including feasibility assessment, eligibility criteria screening, data elements extraction, and text data analytics. As a result, an increasing number of clinical research benefits from state-of-the-art NLP solutions and have been reported ranging from disease study areas15, 16, 17, 18 to drug-related studies19, 20. A majority of existing clinical NLP studies are, however, done within a mono- institutional environment13, which may suffer from limited external validity and research inclusiveness. Compared with single-site research, multisite research potentially offers larger sample size, more adequate representation of participant demographics (e.g., age, gender, race, ethnicity, and social-economic status), and more diverse investigator expertise, which may ultimately yield a higher level of research evidence21, 22, 23, 24.
Despite a plethora of recent advances in adopting NLP for clinical research, there have been barriers towards adoption of NLP solutions in clinical and translation research, especially in multi- site settings. The root causes of these barriers can be categorized into two major reasons: 1) heterogeneity of ETL (extract, transform, load) processes between differing sites with their own disparate EHR environments, and 2) human factor variation in gold standard corpus development processes.
ETL Process Heterogeneity. The challenges faced by NLP development and evaluation to facilitate the secondary use of EHR data originate from the complex, voluminous, and dynamic nature of the data being documented and stored within a heterogeneous set of disparate, institution specific, EHR implementations. Variations in EHR system vendors, data infrastructure (e.g., unified, ontology driven, and de-centralized), and institutions’ modes of operation can lead to idiosyncratic ways of clinical documentation, transformed, and representation25. Collecting these data would require a significant expenditure of effort to locate, retrieve, and link EHR data into a specific format26. This variability in ETL processes required to support a high level of data heterogeneity brings additional challenges in the adoption of NLP for clinical and translational research, which substantially limits both the cross-institutional interoperability of developed NLP solutions and the reproducibility of the associated evaluations.
Human factor variation in gold standard corpus development process. The process of developing, evaluating, and deploying NLP solutions in both mono- and multi-site environments can be task-specific, iterative, and complex, often involving a multitude of stakeholders with diverse backgrounds13, 26. A key step prior to model development is corpus annotation, the process of developing a gold standard by marking the occurrence of both task-defined sets of clinical information as well as their associated interpretative linguistic features (e.g., certainty, status) within text documents. Due to the complexity of clinical language, creating such gold standard corpora requires significant expenditure of domain expertise and time as clinical experts regularly make decisions directly affecting study cohort, annotation guideline, and task definitions. Studies have discovered potential biases in clinical decision making and interpretation of clinical guidelines27, in coding of clinical terminologies28, and in interpretation of imaging findings29. This issue can be further exacerbated when conducting multi-site collaborations due to inter-site variations in care practice30, 31, ultimately affecting the validity and reliability of the resulting gold standard corpus. A coordinated, transparent, and collaborative platform is therefore needed to promote open team science collaboration in NLP algorithm development and evaluation through consensus building, process coordination, and best practice sharing.
Built upon our previous work32, 33, here, we proposed an open NLP development framework to address the aforementioned issues through the following components: 1) an interoperable NLP infrastructure for incorporation of different NLP engines utilizing a clinical common data model for data source interfacing and representation with the aim of reducing the impact of the heterogeneity of ETL processes; 2) a transparent multi-site participation workflow on corpus development and evaluation with the aim of addressing the variation in data abstraction and annotation processes between sites; and 3) a user-centric crowdsourcing interface for collaborative ruleset development that enables effectively and efficiently gathering, synthesizing, and fusing site-specific knowledge and findings. To demonstrate the viability of the framework, we conducted a case study where we developed, evaluated, and implemented an NLP algorithm for extracting 34, 35, 36 COVID-19 signs and symptoms to support the National COVID Cohort Collaborative (N3C).
Framework Description
The framework itself consists of a data ingestion layer, a processing layer, and a data persistence layer. The architecture of the proposed framework is illustrated in Figure 3. The data ingestion layer works as the data collector with the ability to read text from a configurable variety of data sources such as relational databases or file systems including and load them into the NOTE table of OMOP CDM. The processing layer serves as the NLP engine where information extraction from raw texts happens given a set of heuristic rules created by various NLP engines. By default, as an example implementation, the MedTagger37 NLP engine is provided, although alternative NLP engines can be substituted by wrapping their respective NLP pipelines to conform to a provided API specification. After the term modifiers added by contextual rules from ConText Algorithm38 around the extracted condition mentions, these conditions will compose clinical events with temporal information. The reason we opt for a symbolic solution is due to its simplicity, transparency, and interpretability as the outcomes are fully deterministic based on the definition of the rules. When the baseline rulesets and dictionaries are made available to the public, they can therefore be easily refined by different users from different sites. The data persistence layer stores resulting extracted NLP artifacts in the OMOP CDM NOTE_NLP table as the events are extracted from NLP systems.
The framework is distributed as open-source software under the Apache 2.0 license via Github in three parts: 1) ETL Backbone (https://github.com/OHNLP/Backbone) with an example NLP
Documentation
(https://github.com/OHNLP/N3C-NLP-Documentation/wiki),
Collaborative
platform for developing NLP rulesets (https://github.com/OHNLP/OHNLPTK). The demo homepage (Figure 2(a) - https://ohnlp4covid-dev.n3c.ncats.io/) demonstrates the N3C NLP engine outputs on annotating clinical text using the baseline rulesets and dictionary. The annotations are from components of Sign/Symptom extractor, temporal information extractor and dictionary lookup extractor. To further customize each model, the users can visit “Rule Editor”
Builder”
(https://ohnlp4covid-dev.n3c.ncats.io/dict_builder) page (Figure 2(b)). Figure 2(c) provides an example of the rules editing interface with the baseline COVID-19 ruleset. The rulesets can be tested in real time by clicking the “Upload and test” button, where the rulesets will be uploaded, and the NLP engine will be generated for testing and debugging purposes. As a use case study, we also provide an example NLP project for extracting signs/symptoms related to COVID-19 that was developed as an example use case for this framework. The elements with original texts such as text snippets and concept mentions are truncated before submission.
N3C Case Study
NLP Algorithm Development and Evaluation: Table 1 shows the annotation corpora statistics. A COVID-19 sign/symptom ruleset was produced consisting of 17 concepts. The IAA of the annotated corpus was 0.686 F1-score for Mayo, 0.516 for UMN and 0.211 for UKen. Two NLP algorithms were evaluated in this study. One was developed based solely on the narratives sourced from a single site (Mayo Clinic). The other used the resulting NLP algorithm from the single site and fine-tuned based on the annotated training data from an additional two sites (UMN and UKen).
Table 2 shows the performance of the single-site NLP algorithm and Table 3 shows the performance of the multi-site NLP algorithm. The single-site ruleset resulted in performances of 0.876, 0.706, and 0.694 in F-scores for Mayo, Minnesota, and Kentucky test datasets, respectively, while the multi-site NLP ruleset improved performances to 0.884, 0.769 and 0.806. The performance of the multi-site NLP algorithm was better than that of the single-site NLP algorithm, but both showed a degrading trend from Mayo site to other sites.
Tables 4, 5 and 6 show the results of error analysis for the three sites. For FP, major discrepancies between the NLP algorithm and the gold standard were due to the NLP algorithm extracting mentions that are not COVID signs/symptoms but for instruction/patient education, adverse events/indication of treatment, clinical goal/precaution, template, etc. It should be noted that gold standards were not always correct, and in some notes, it was hard to judge if the mentions are COVID signs/symptoms when symptoms are not appearing with COVID or de-identified dates are inconsistent. For FN, reasons include NLP algorithm not complete, tokenization error due to de- identification process, template, and annotation errors.
Discussion
In this study, we proposed an open NLP development framework with the following properties: an interoperable NLP infrastructure, a transparent multi-site participation workflow, and a user- centric crowdsourcing interface. The key goal of this framework is to facilitate multi-site collaborative development, evaluation, and implementation of NLP algorithms. The framework has been implemented to support efforts conducted by the National COVID Cohort Collaborative (N3C) to enable the utilization of unstructured text in high throughput.
Here, we have presented our results from running our framework using a centralized annotation process on texts sourced from multiple sites after de-identification, with the aim of assessing the impact on NLP algorithm development (single-site algorithm vs multi-site algorithm). Several pragmatic implementation challenges were discovered that may impact the intermediate and final NLP results. We observed that IAA varied greatly between the three sites despite the fact that annotators had been trained using de-identified Mayo notes (0.686 F1-score for Mayo, 0.516 for UMN, 0.211 for UKen). Firstly, utilizing a centralized annotation approach, the process of text data collection took a very long time because each site needs to complete de-identification before sharing data. Secondly, it was a challenge for annotators to work on annotation tasks that spanned a long period of time. Thirdly, the shared data sets were usually small, and as such, annotators had no chance to do annotation training using these outside notes, and it was hard for them to get familiar with the disparate variety of document structures from other sites.
Both multi-site and single-site NLP algorithms showed a degrading trend in performance from Mayo site to other sites, albeit this issue being less prominent in the multi-site NLP algorithm as compared to the single-site NLP algorithm. The data sharing issues also impacted NLP algorithm performance. First, training sets from outside institutions were very small due to the small number of shared notes, causing difficulties in developing comprehensive rules as features, patterns, and contextual information that could appear in third party narratives could not be fully represented in such a small sample. Second, de-identification processes could cause text span issues that may impact the input text format and thus NLP algorithm performance. The algorithm performance for algorithms developed through a centralized mode was therefore not ideal for immediate use at multiple sites, as additional local fine-tuning is still needed before final implementation and application.
Our experiment results showed that a centralized approach towards multi-site NLP algorithm development is suboptimal for advancing the adoption of NLP techniques in the clinical and translational research community, this further support our proposed federated method. The experiment also demonstrates that deployment of NLP algorithms for multi-site studies needs to be done in each local site. To ensure the scientific rigor of the data generated, each site need to perform annotation and evaluation on their own while collectively contributing to NLP algorithm development and refinement. Since the NLP models are to be shared in rule-based systems, the models can be shared without the concerns typically associated with language resources involving the Protected Health Information (PHI) issue.
In the proposed workflow, each site will evaluate the NLP algorithms for concept extraction by creating a gold standard corpus based on the common annotation guidelines. The federated evaluation can be deployed leveraging cloud computing through a centralized controller where NLP algorithms can be distributed to each institution. NLP Sandbox1 is an example of such an evaluation framework, which uses Docker39 containers to encapsulate algorithm implementations.
By adopting this process, the evaluation only happens behind each institution’s firewall, and only the summary statistics on NLP algorithm performance (i.e., no raw data containing PHI) is transferred out of the firewall. Performance statistics, such as the precision, recall, and F1-score, as defined depending on the experimental setting, can be obtained in near real-time and can thus be used as part of continuous development workflows.
This federated process offers several benefits. For instance, when conducting error analysis, we discovered that contexts played an important role in this case study. Error analyses showed it was not a trivial task to extract COVID signs/symptoms, as their occurrence is not necessarily isolated only to occurring due to COVID, and as they could appear as adverse events/indication of treatment, or in instruction/patient education, or clinical goal/precaution, etc. This posed a challenge not only for annotation, but also for the NLP algorithm development. One benefit of the
1 Nlp Sandbox: Https://Github.Com/Nlpsandbox
federated annotation and development process is that these contexts can be systematically incorporated by local expertise in the annotation process. Deployment of a federated development framework requires the participation of multiple sites.
Adoption can, however, be hindered by the fact that the process of translating NLP algorithms into implementation is complex, much like the “bench to bedside” process that translates laboratory discoveries into patient care. To facilitate participation in our federated method, we have developed a further suite of tools such as MedTator40 and best practice guidelines41. MedTator, a serverless annotation tool, aims to provide an intuitive and interactive user interface for high- quality annotation corpus generation. The best practice guideline contains detailed instructions for facilitating multisite annotation practice with the following key activities: task formulation, cohort screening, annotation guideline development, annotation training, annotation production, and adjudication.
Simply having the toolsets be available, is, however, insufficient. Pragmatically, we have seen that there is a hyper focus on novel methods in academia with competing as opposed to collaborative priorities in NLP algorithm development. Our experience suggests that a collaborative development process for NLP algorithms is needed for truly implementable and useful multi-site NLP solutions. This is one of the key goals we seek to achieve with the Open Health Natural Language Processing (OHNLP) Collaboratory and have thus positioned our framework’s workflow to facilitate this task. Additionally, we recognize that it is not simply a software problem, a local workforce is also needed at each institution. As a consequence of conducting coordinated development of NLP algorithms deployed using our framework as a solution for consortia-specific tasks such as with the N3C, we simultaneously build the human workforce locally at institutions necessary to conduct the federated development, evaluation, and implementation of NLP algorithms using our framework.
Design Principles
Incorporating standards and interoperability. A common barrier to the widespread adoption of NLP in clinical research is the need to transform input and outputs to conform to part of an overall pipeline. While seemingly straightforward, such a task is difficult without prior significant investment in associated infrastructure and dedicated software development. It is therefore desirable to leverage existing infrastructure where possible and incorporate such an effort into the distributed NLP pipeline to reduce technical burden on the end user.
There is, however, significant variation in terms of available infrastructure and data availability amongst different institutions. Creating a solution that is immediately suitable for all these environments out of the box would be immensely challenging. For that reason, we sought to leverage existing data modeling efforts that are likely to be already adopted by academic medical institutions to standardize the data ingestion and output process. In our implementation, we chose the Observational Health Data Sciences and Informatics’ Observational Medical Outcomes Partnership common data model (OHDSI/OMOP CDM) to handle input of clinical narratives via the NOTE table and output via the NOTE_NLP table. This brings the advantage that input/output is now standardized: so long as institutions have already transformed their clinical data into the OMOP CDM, and/or their downstream NLP-reliant applications read from the OMOP CDM database, no additional technical development burden is needed.
It is important to note that standardization as a default only serves to simplify adoption for those who already have a solution complying with the standard and cannot be a comprehensive solution. A purely OMOP CDM reliant solution is not ideal, as not all institutions will have their own OMOP CDM instance and standing up such an instance to just use a pipeline may produce undue burden.
For that reason, input/output in our infrastructure is modularized, and can be substituted at will: the default OMOP CDM I/O utilizes a variant of SQL-based data extractors/writers, and the specific query and connection strings used can be substituted via plaintext configuration changes.
Additionally, SQL-based I/O is not the only supported setting, a variety of other data sources including Elasticsearch, google cloud storage, amazon s3, and plaintext are included as well as configuration-swappable options.
Crowdsourcing algorithm development. To promote collaboration and sharing efforts between participants in the algorithm development process, we built a crowdsourcing platform for domain experts to upload, customize, and examine their NLP algorithms in an interactive web application.
Users can create keyword-based and rule-based algorithms and test the performance in the online environment instantly. The crowdsourcing platform consists of three modules based on our NLP system to support expert collaboration, including dictionary builder, regular expression rule set editor, and detection result visualization.
The dictionary builder can extend the keyword collection used by the algorithm. Users can customize particular terms from the ontology database such as CIDO 42 and MONDO 43. The regular expression rule set editor provides an integrated interface to help users customize their own regular expression rule set (on top of an existing dictionary, if desired), to support use cases such as extraction of new symptoms, treatments, or outcomes. The detection result visualization is designed based on Brat annotation tool 44 to check the results generated by different methods.
Case Study
The National COVID Cohort Collaborative36 (N3C) is a novel partnership that includes the Clinical and Translational Science Awards (CTSA) Program hubs, the National Center for Advancing Translational Science (NCATS), the Center for Data to Health (CD2H) and the community, focusing on collaborative sharing of structured EHR data. Access to unstructured data is limited due to protection of PHI and clinical care decision logics, that were further contributing to NLP infrastructure lacking within the consortia. However, structured data does not show the whole picture from the EHR perspective, greatly restricting research activities. In this case study, extraction of COVID-19 signs and symptoms was used as a case study to investigate the viability of the proposed framework among sites participating in the N3C.
Centralizing gold standard corpus development. Due to resources and time constraints at each of the N3C sites, we opted to conduct the gold standard corpus development process in a centralized manner. A collection of de-identified and synthesized clinical documents was gathered from participating sites through an existing de-identification effort led by the NCATS Clinical Data to Health (CD2H). The N3C deidentification and synthetic text generation workflow is illustrated in Figure 1. Specifically, clinical notes from patients with positive COVID-19 test visit notes (e.g., nurse calls, etc.), notes that had fewer than 1000 characters, and notes that were authored more than 14 days prior to the date of the patient’s earliest positive COVID-19 test result were further filtered out. A total of 369 clinical notes from these sites that met these criteria were randomly selected, de-identified using the de-identification program developed by the Medical College of Wisconsin followed by manual review. The removed PHI identifiers are replaced by the programmatically added synthetic texts. We collected 20 signs and symptoms of COVID-19 as a basic COVID-19 concept set according to the recommendations from the CDC and Mayo Clinic. Five out of the 20 concepts are emergency warning signs including dyspnea, chest pain, delirium, hypersomnia and cyanosis. We then gathered formal definitions of each clinical concept from the Coronavirus Infectious Disease Ontology (CIDO) 42. Based on the Open Biological and Biomedical Ontology (OBO) Foundry library, CIDO concepts were imported from 45 ontologies, and it uses Human Phenotype Ontology (HPO) 45 for phenotypes. Some representative phenotypes shown in COVID-19 have been imported to the CIDO. However, if the chosen COVID-19 clinical concepts were not collected by the CIDO, we re-pulled them from the HPO to the CIDO. We also gathered cross-reference concept codes from the CIDO including UMLS 46, SNOMED-CT 47, MeSH 48, HPO, MeDDRA 49.
We selected available clinical notes from both inpatients and outpatients in the two-week window preceding the order date of the first positive COVID-19 result as the annotation cohort. After the text data was collected from participating sites, the same annotation process was completed by the annotator team from Mayo Clinic to generate the gold standard annotations on COVID-19 signs and symptoms. There are 313 clinical notes from Mayo Clinic, 20 notes from UKen and 36 notes from UMN. Annotators were first trained using Mayo notes to gain better understanding of the annotation guidelines. Inter-annotator agreement (IAA) was calculated after annotation and corresponding discrepancies were resolved by discussions between the two annotators to generate a final gold standard dataset.
NLP algorithm development and evaluation. Using the annotated corpus, we developed both a single-site and multi-site NLP algorithm using a regular expression-based matching method, which has been widely adopted for information extraction in clinical settings. Specifically, for the Mayo data, we randomly chose 101 notes out of the 313 annotated notes as development set, 105 notes as validation set, and the remaining 107 notes were used as the testing set. For the UKen data, 10 notes were used for training and 10 for testing. For the UMN data, 18 was used for training and 18 for testing. Single-site algorithm was developed using the development set and validation set from Mayo, tested on the Mayo testing set and all data from UKen and UMN. Multi-site algorithm was generated through further refinement of the single-site algorithm using training sets from UKen and UMN and then tested on testing sets from all sites.
We evaluated the performance of single-site and multi-site algorithms using precision, recall, and F1-score for the annotated concept mentions, without and with certainty. A span can be represented from the start position to the end position of the concept mention. Certainty is an attribute of the concept mention including positive, negated, hypothetical and possible. For the mention-level evaluation without certainty, when there are overlaps between the gold standard mention span and the NLP detected mention span while the concept type (i.e., the specific sign/symptom such as fever, cough) is the same, it is considered a true positive (TP). If a concept mention exists in the gold standard annotation but not detected by the NLP algorithm, or spans overlap but the concept type is not matched, it is considered as a false negative (FN). If a concept mention is detected by the algorithm but does not exist in the gold standard annotation, the concept is considered as a false positive (FP). For the mention-level span and certainty evaluation, certainty match needs to be considered when calculating TP, FN and FP. The precision, recall and F1-score are then calculated as follows. We further manually analyzed errors from multi-site algorithm mention- level evaluation without certainty.
Acknowledgment
This research was possible because of the patients whose information is included within the data and the organizations and scientists who have contributed to the on-going development of this community resource https://doi.org/10.1093/jamia/ocaa196. The analyses described in this publication were conducted with data or tools accessed through the NCATS N3C Data Enclave https://covid.cd2h.org and N3C Attribution & Publication Policy v1.2-2020-08-25b and supported by NCATS U24 TR002306. This task was made possible by the National Center for Advancing Translational Sciences of the National Institutes of Health under award number U01TR02062 and the Bill & Melinda Gates Foundation. The content is solely the responsibility of the authors and does not necessarily represent the official views of the National Institutes of Health.
• Etl Pipeline: Https://Github.Com/Ohnlp/Backbone
• NLP Implementation: https://github.com/OHNLP/MedTagger • Web Rule Editor Front-end: https://github.com/OHNLP/OHNLPTK • MedTator annotation tool: https://github.com/OHNLP/MedTator
The Developed Nlp Ruleset Can Be Found At
https://github.com/OHNLP/covid19ruleset/tree/main/covid19
Ai Adaptive Learning
This project focuses on ai adaptive learning using modern AI and machine learning techniques. The content below is adapted from research literature and practical implementation notes.
We propose a novel high-performance and interpretable canon-
addition, unlike tree learning, DNNs enable gradient descent- ical deep tabular data learning architecture, TabNet. TabNet based end-to-end learning for tabular data which can have a uses sequential attention to choose which features to reason multitude of benefits: (i) efficiently encoding multiple data from at each decision step, enabling interpretability and more types like images along with tabular data; (ii) alleviating the efficient learning as the learning capacity is used for the most need for feature engineering, which is currently a key aspect
salient features. We demonstrate that TabNet outperforms in tree-based tabular data learning methods; (iii) learning other variants on a wide range of non-performance-saturated from streaming data and perhaps most importantly (iv) end- tabular datasets and yields interpretable feature attributions to-end models allow representation learning which enables plus insights into its global behavior. Finally, we demonstrate many valuable application scenarios including data-efficient
self-supervised learning for tabular data, significantly improv- domain adaptation (Goodfellow, Bengio, and Courville 2016), ing performance when unlabeled data is abundant. generative modeling (Radford, Metz, and Chintala 2015) and
Introduction We propose a new canonical DNN architecture for tabular
Deep neural networks (DNNs) have shown notable success data, TabNet. The main contributions are summarized as: efficiently encode the raw data into meaningful representa- enabling flexible integration into end-to-end learning. tions, fuel the rapid progress. One data type that has yet to 2. TabNet uses sequential attention to choose which fea- see such success with a canonical architecture is tabular data. tures to reason from at each decision step, enabling in-
Despite being the most common data type in real-world AI terpretability and better learning as the learning capacity (as it is comprised of any categorical and numerical features), is used for the most salient features (see Fig. 1). This under-explored, with variants of ensemble decision trees for each input, and unlike other instance-wise feature se- Why? First, because DT-based approaches have certain bene- and van der Schaar 2019), TabNet employs a single deep
fits: (i) they are representionally efficient for decision mani- learning architecture for feature selection and reasoning. folds with approximately hyperplane boundaries which are 3. Above design choices lead to two valuable properties: (i) common in tabular data; and (ii) they are highly interpretable TabNet outperforms or is on par with other tabular learn- in their basic form (e.g. by tracking decision nodes) and there ing models on various datasets for classification and re-
are popular post-hoc explainability methods for their ensem- gression problems from different domains; and (ii) TabNet ble form, e.g. (Lundberg, Erion, and Lee 2018) – this is an enables two kinds of interpretability: local interpretability important concern in many real-world applications; (iii) they that visualizes the importance of features and how they are fast to train. Second, because previously-proposed DNN are combined, and global interpretability which quantifies
architectures are not well-suited for tabular data: e.g. stacked the contribution of each feature to the trained model. convolutional layers or multi-layer perceptrons (MLPs) are 4. Finally, for the first time for tabular data, we show signif- vastly overparametrized – the lack of appropriate inductive icant performance improvements by using unsupervised bias often causes them to fail to find optimal solutions for tab- pre-training to predict masked features (see Fig. 2).
ular decision manifolds (Goodfellow, Bengio, and Courville
Why is deep learning worth exploring for tabular data?
One obvious motivation is expected performance improve- Feature selection: Feature selection broadly refers to judi- Copyright © 2021, Association for the Advancement of Artificial ciously picking a subset of features based on their useful-
Professional occupation related Investment related
Feedback from Feedback to
Feature selection Input processing Feature selection Input processing
previous step next step … …
Predicted output (whether the income level >$50k)
selection enables interpretability and better learning as the capacity is used for the most salient features. TabNet employs multiple decision blocks that focus on processing a subset of input features for reasoning. Two decision blocks shown as examples process features that are related to professional occupation and investments, respectively, in order to predict the income level.
Unsupervised pre-training Supervised fine-tuning
Age Cap. gain Education Occupation Gender Relationship Age Cap. gain Education Occupation Gender Relationship 5 2000 ? Exec-managerial F Wife 6 2000 Bachelors Exec-managerial M Husband 1 0 ? Farming-fishing M ? 2 0 High-school Farming-fishing M Unmarried
? 50 Doctorate Prof-specialty M Husband 4 50 Doctorate Prof-specialty M Husband 2 ? ? Handlers-cleaners F Wife 2 0 High-school Handlers-cleaners F Wife 5 3000 Bachelors ? ? Husband 5 3000 Bachelors Exec-managerial M Husband
3 0 Bachelors ? F ? 3 100 Bachelors Prof-specialty F Wife ? 0 High-school Armed-Forces ? Husband 2 0 High-school Armed-Forces M Husband
TabNet decoder Decision making
Age Cap. gain Education Occupation Gender Relationship Income > $50k
3 M False
level can be guessed from the occupation, or the gender can be guessed from the relationship. Unsupervised representation learning by masked self-supervised learning results in an improved encoder model for the supervised learning task.
ward selection and Lasso regularization (Guyon and Elisseeff performance with compact representations. 2003) attribute feature importance based on the entire training Tree-based learning: DTs are commonly-used for tabular data, and are referred as global methods. Instance-wise fea- data learning. Their prominent strength is efficient picking ture selection refers to picking features individually for each of global features with the most statistical information gain
to maximize the mutual information between the selected mance of standard DTs, one common approach is ensembling features and the response variable, and in (Yoon, Jordon, and to reduce variance. Among ensembling methods, random van der Schaar 2019) by using an actor-critic framework to forests (Ho 1998) use random subsets of data with randomly mimic a baseline while optimizing the selection. Unlike these, selected features to grow many trees. XGBoost (Chen and
sity in end-to-end learning – a single model jointly performs recent ensemble DT approaches that dominate most of the feature selection and output mapping, resulting in superior recent data science competitions. Our experimental results
!# + Softmax !" < % !" > % !# > & !# > &
ReLU ReLU &
$" !" − $" % −1 −$" !" + $" % −1 −1 $# !# − $# & % −1 −$# !# + $# & !"
FC FC
W: [$" , - $" , 0, 0] W: [0, 0, $# , - $# ] !" < % b: [-a $" , a $" , -1, -1] b: [-1, -1, -d $# , d $# ] !# < & !" > % !# < & [!" ] [!# ]
M: [1, 0] M: [0, 1]
(right). Relevant features are selected by using multiplicative sparse masks on inputs. The selected features are linearly transformed, and after a bias addition (to represent boundaries) ReLU performs region selection by zeroing the regions. Aggregation of multiple regions is based on addition. As C and C get larger, the decision boundary gets sharper.
for various datasets show that tree-based models can be out- constructs a sequential multi-step architecture, where each performed when the representation capacity is improved with step contributes to a portion of the decision based on the deep learning while retaining their feature selecting property. selected features; (iii) improves the learning capacity via non- Integration of DNNs into DTs: Representing DTs with linear processing of the selected features; and (iv) mimics
DNN building blocks as in (Humbird, Peterson, and McClar- ensembling via higher dimensions and more steps. ren 2018) yields redundancy in representation and ineffi- cient learning. Soft (neural) DTs (Wang, Aggarwal, and Liu Fig. 4 shows the TabNet architecture for encoding tabu- functions, instead of non-differentiable axis-aligned splits. mapping of categorical features with trainable embeddings.
However, losing automatic feature selection often degrades We do not consider any global feature normalization, but performance. In (Yang, Morillo, and Hospedales 2018), a soft merely apply batch normalization (BN). We pass the same D- binning function is proposed to simulate DTs in DNNs, by dimensional features f ∈ <B×D to each decision step, where 2019) proposes a DNN architecture by explicitly leveraging multi-step processing with Nsteps decision steps. The ith
expressive feature combinations, however, learning is based step inputs the processed information from the (i − 1)th step on transferring knowledge from gradient-boosted DT. (Tanno to decide which features to use and outputs the processed ing from primitive blocks while representation learning into sion. The idea of top-down attention in the sequential form edges, routing functions and leaf nodes. TabNet differs from is inspired by its applications in processing visual and text
these as it embeds soft feature selection with controllable data (Hudson and Manning 2018) and reinforcement learn- Self-supervised learning: Unsupervised representation relevant information in high dimensional input. learning improves supervised learning especially in small Feature selection: We employ a learnable mask M[i] ∈ has shown significant advances – driven by the judicious capacity of a decision step is not wasted on irrelevant
choice of the unsupervised learning objective (masked input ones, and thus the model becomes more parameter effi- prediction) and attention-based deep learning. cient. The masking is multiplicative, M[i] · f . We use an attentive transformer (see Fig. 4) to obtain the masks us- TabNet for Tabular Learning ing the processed features from the preceding step, a[i − 1]:
M[i] = sparsemax(P[i − 1] · hi (a[i − 1])). Sparsemax nor-
DTs are successful for learning from real-world tabular malization (Martins and Astudillo 2016) encourages sparsity datasets. With a specific design, conventional DNN building by mapping the Euclidean projection onto the probabilistic blocks can be used to implement DT-like output manifold, simplex, which is observed to be superior in performance and e.g. see Fig. 3). In such a design, individual feature selec- aligned with the goal of sparse feature selection for explain-
tion is key to obtain decision boundaries in hyperplane form, PD which can be generalized to a linear combination of features ability. Note that j=1 M[i]b,j = 1. hi is a trainable func- where coefficients determine the proportion of each feature. tion, shown in Fig. 4 using a FC layer, followed by BN. P[i] TabNet is based on such functionality and it outperforms DTs is the prior scale term, denoting how much a particular feature
Qi while reaping their benefits by careful design which: (i) uses has been used previously: P[i] = j=1 (γ − M[j]), where γ sparse instance-wise feature selection learned from data; (ii) is a relaxation parameter – when γ = 1, a feature is enforced
+ Softmax
Feature Feature …
transformer transformer
x Nsteps Features
+ Softmax
Feature …
transformer transformer Feature Feature Feature Feature transformer
Encoded representation
transformer transformer Attentive transformer … Mask transformer …
Step 2 Decision step dependent
transformer transformer
BN Feature Feature
FC BN transformer transformer
+ 0.5 0.5 0.5 Agg. Agg. Features Features FC FC + +
Reconstructed + … Feature attributes + … features
(a) TabNet encoder architecture (b) TabNet decoder architecture Feature transformer Feature Attentive transformer Shared across decision steps Decision step dependent transformer GLU
Decision step dependent Prior scales
+ 0.5 0.5 0.5
0.5 0.5 0.5
+ Attentive transformer (c) (d)
Prior scales
divides the processed representation to be used by the attentive transformer of the subsequent step as well as for the overall Attentive BN FC
output. For each step, the feature selection mask provides interpretable information about the model’s functionality, and the +
masks can be aggregated to obtain global feature transformer important attribution. (b) TabNet decoder, composed of a feature transformer block at each step. (c) A feature transformer block example – 4-layer network is shown, where 2 are shared across all decision
Prior scales
steps and 2 are decision step-dependent. Each layer is composed of a fully-connected (FC) layer, BN and GLU nonlinearity. (d) +
An attentive transformer block example – a single layer mapping is modulated with a prior scale information which aggregates Sparsemax
how much each feature has been used before the current decision step. sparsemax (Martins and Astudillo 2016) is used for BN FC
normalization of the coefficients, resulting in sparse selection of the salient features. +
to be used only at one decision step and as γ increases, more propose the aggregate.feature importance mask, Magg−b,j = flexibility is provided to use a feature at multiple decision PNsteps ηb [i]Mb,j [i]
PD PNsteps
ηb [i]Mb,j [i].2 i=1 i=1 steps. P is initialized as all ones, 1B×D , without any prior j=1
on the masked features. If some features are unused (as in self- Tabular self-supervised learning: We propose a decoder supervised learning), corresponding P entries are made 0 architecture to reconstruct tabular features from the Tab- to help model’s learning. To further control the sparsity of the Net encoded representations. The decoder is composed of selected features, we propose sparsity regularization in the feature transformers, followed by FC layers at each deci-
form of entropy (Grandvalet and Bengio 2004), Lsparse = sion step. The outputs are summed to obtain the recon-
PNsteps PB PD −Mb,j [i] log(Mb,j [i]+)
i=1 b=1 j=1 Nsteps ·B , where is a structed features. We propose the task of prediction of miss- small number for numerical stability. We add the sparsity reg- ing feature columns from the others. Consider a binary mask ularization to the overall loss, with a coefficient λsparse . Spar- S ∈ {0, 1}B×D . The TabNet encoder inputs (1 − S) · f̂ sity provides a favorable inductive bias for datasets where and the decoder outputs the reconstructed features, S · f̂ . We
most features are redundant. initialize P = (1 − S) in encoder so that the model em- Feature processing: We process the filtered features using phasizes merely on the known features, and the decoder’s last a feature transformer (see Fig. 4) and then split for the FC layer is multiplied with S to output the unknown features. decision step output and information for the subsequent We consider the reconstruction loss in self-supervised phase:
step, [d[i], a[i]] = fi (M[i] · f ), where d[i] ∈ <B×Nd and 2
PB PD (f̂b,j −fb,j )·Sb,j
a[i] ∈ <B×Na . For parameter-efficient and robust learning b=1 j=1
√ PB PB 2
. Normalization b=1 (fb,j −1/B b=1 fb,j ) with high capacity, a feature transformer should comprise layers that are shared across all decision steps (as the same with the population standard deviation of the ground truth features are input across different decision steps), as well as is beneficial, as the features may have different ranges. We decision step-dependent layers. Fig. 4 shows the implementa- sample Sb,j independently from a Bernoulli distribution with
tion as concatenation of two shared layers and two decision parameter ps , at each iteration. step-dependent layers. Each FC layer is followed by BN and eventually connected to a normalized residual √ connection We study TabNet in wide range of problems, that contain with normalization. Normalization with 0.5 helps to sta- regression or classification tasks, particularly with published bilize learning by ensuring that the variance throughout the benchmarks. For all datasets, categorical inputs are mapped
For faster training, we use large batch sizes with BN. Thus, bedding and numerical columns are input without and pre- except the one applied to the input features, we use ghost BN processing.4 We use standard classification (softmax cross (Hoffer, Hubara, and Soudry 2017) form, using a virtual batch entropy) and regression (mean squared error) loss functions size BV and momentum mB . For the input features, we ob- and we train until convergence. Hyperparameters of the Tab-
serve the benefit of low-variance averaging and hence avoid Net models are optimized on a validation set and listed in ghost BN. Finally, inspired by decision-tree like aggregation Appendix. TabNet performance is not very sensitive to most as in Fig. 3, we construct the overall decision embedding hyperparameters as shown with ablation studies in Appendix. as dout = i=1 PNsteps ReLU(d[i]). We apply a linear mapping In Appendix, we also present ablation studies on various de-
Wfinal dout to get the output mapping.1 sign and guidelines on selection of the key hyperparameters. Interpretability: TabNet’s feature selection masks can shed For all experiments we cite, we use the same training, val- light on the selected features at each step. If Mb,j [i] = 0, idation and testing data split with the original work. Adam optimization algorithm (Kingma and Ba 2014) and Glorot then j th feature of the bth sample should have no contribution uniform initialization are used for training of all models.5
to the decision. If fi were a linear function, the coefficient
Mb,j [i] would correspond to the feature importance of fb,j . Instance-wise feature selection
Although each decision step employs non-linear processing, their outputs are combined later in a linear way. We aim Selection of the salient features is crucial for high perfor- to quantify an aggregate feature importance in addition to mance, especially for small datasets. We consider 6 tabular requires a coefficient that can weigh the relative importance samples). The datasets are constructed in such a way that of each step in the decision. We simply propose ηb [i] = only a subset of the features determine the output. For Syn1-
PNd Syn3, salient features are same for all instances (e.g., the
c=1 ReLU(db,c [i]) to denote the aggregate decision con- tribution at ith decision step for the bth sample. Intuitively, if 2
Normalization is used to ensure D
P j=1 Magg−b,j = 1. db,c [i] < 0, then all features at ith decision step should have 3
0 contribution to the overall decision. As its value increases, prove the performance, but interpretation of individual dimensions
it plays a higher role in the overall linear combination. Scal- may become challenging. ing the decision mask at each decision step with ηb [i], we Specially-designed feature engineering, e.g. logarithmic trans- formation of variables highly-skewed distributions, may further
For discrete outputs, we additionally employ softmax during
training (and argmax during inference). An open-source implementation will be released.
Global: using only globally-salient features, Tree Ensembles (Geurts, Ernst, and Wehenkel 2006), Lasso-regularized model, L2X
Syn Syn Syn Syn Syn Syn
No selection .5 ± .0 .7 ± .0 .8 ± .0 .5 ± .0 .6 ± .0 .6 ± .0 Tree .5 ± .1 .8 ± .0 .8 ± .0 .6 ± .0 .7 ± .0 .7 ± .0 Lasso-regularized .4 ± .0 .5 ± .0 .8 ± .0 .5 ± .0 .6 ± .0 .7 ± .0
INVASE .6 ± .0 .8 ± .0 .9 ± .0 .7 ± .0 .7 ± .0 .8 ± .0
Global .6 ± .0 .8 ± .0 .9 ± .0 .7 ± .0 .7 ± .0 .8 ± .0 TabNet .6 ± .0 .8 ± .0 .8 ± .0 .7 ± .0 .7 ± .0 .8 ± .0
output of Syn depends on features X -X ), and global fea- Table 3: Performance for Poker Hand induction dataset. ture selection, as if the salient features were known, would give high performance. For Syn4-Syn6, salient features are Model Test accuracy (%) instance dependent (e.g., for Syn4, the output depends on ei- DT 50.0 ther X -X or X -X depending on the value of X ), which MLP 50.0
makes global feature selection suboptimal. Table 1 shows that Deep neural DT 65.1
TabNet outperforms others (Tree Ensembles (Geurts, Ernst, XGBoost 71.1
and Wehenkel 2006), LASSO regularization, L2X (Chen LightGBM 70.0 van der Schaar 2019). For Syn1-Syn3, TabNet performance TabNet 99.2 is close to global feature selection - it can figure out what Rule-based 100.0 features are globally important. For Syn4-Syn6, eliminating instance-wise redundant features, TabNet improves global feature selection. All other methods utilize a predictive model Poker Hand (Dua and Graff 2017): The task is classifica-
with 43k parameters, and the total number of parameters is tion of the poker hand from the raw suit and rank attributes of 101k for INVASE due to the two other models in the actor- the cards. The input-output relationship is deterministic and critic framework. TabNet is a single architecture, and its size hand-crafted rules can get 100% accuracy. Yet, conventional is 26k for Syn1-Syn and 31k for Syn4-Syn6. The compact DNNs, DTs, and even their hybrid variant of deep neural DTs
representation is one of TabNet’s valuable properties. (Yang, Morillo, and Hospedales 2018) severely suffer from the imbalanced data and cannot learn the required sorting and Performance on real-world datasets ranking operations (Yang, Morillo, and Hospedales 2018).
Tuned XGBoost, CatBoost, and LightGBM show very slight
as it can perform highly-nonlinear processing with its depth, Model Test accuracy (%) without overfitting thanks to instance-wise feature selection.
CatBoost 85.1 Table 4: Performance on Sarcos dataset. Three TabNet mod-
AutoML Tables 94.9 els of different sizes are considered.
Forest Cover Type (Dua and Graff 2017): The task is clas- MLP 2.1 0.14M
sification of forest cover type from cartographic variables. Adaptive neural tree 1.2 0.60M approaches that are known to achieve solid performance (AutoML 2019), an automated search framework based on TabNet-M 0.2 0.59M ensemble of models including DNN, gradient boosted DT, TabNet-L 0.1 1.75M with very thorough hyperparameter search. A single TabNet without fine-grained hyperparameter search outperforms it. Sarcos (Vijayakumar and Schaal 2000): The task is re-
gressing inverse dynamics of an anthropomorphic robot arm.
very small model is possible with a random forest. In the very and TabNet merely focuses on the relevant ones. For Syn4, small model size regime, TabNet’s performance is on par the output depends on either X -X or X -X depending parameters. When the model size is not constrained, TabNet feature selection – it allocates a mask to focus on the indi- achieves almost an order of magnitude lower test MSE. cator X , and assigns almost all-zero weights to irrelevant
features (the ones other than two feature groups). models are denoted with -S and -M. Real-world datasets: We first consider the simple task of mushroom edibility prediction (Dua and Graff 2017). Tab- Model Test acc. (%) Model size Net achieves 100% test accuracy on this dataset. It is indeed Sparse evolutionary MLP 78.4 81K known (Dua and Graff 2017) that “Odor” is the most discrim-
What is this project about?
This project covers practical implementation and research aspects of the topic using AI/ML techniques.