Machine Learning In Big Data Platform
Abdelrahim Kasem Ahmad* , Assef Jafar and Kadan Aljoumaa
Introduction
The telecommunications sector has become one of the main industries in developed countries. The technical progress and the increasing number of operators raised the level of competition . Companies are working hard to survive in this competitive market depending on multiple strategies. Three main strategies have been proposed to generate more revenues : (1) acquire new customers, (2) upsell the exist- ing customers, and (3) increase the retention period of customers. However, com- paring these strategies taking the value of return on investment (RoI) of each into account has shown that the third strategy is the most profitable strategy , proves that retaining an existing customer costs much lower than acquiring a new one , in addition to being considered much easier than the upselling strategy . To apply the
Abstract
Customer churn is a major problem and one of the most important concerns for large companies. Due to the direct effect on the revenues of the companies, especially in the telecom field, companies are seeking to develop means to predict potential customer to churn. Therefore, finding factors that increase customer churn is important to take necessary actions to reduce this churn. The main contribution of our work is to develop a churn prediction model which assists telecom operators to predict cus- tomers who are most likely subject to churn. The model developed in this work uses machine learning techniques on big data platform and builds a new way of features’ engineering and selection. In order to measure the performance of the model, the Area Under Curve (AUC) standard measure is adopted, and the AUC value obtained is 93.3%. Another main contribution is to use customer social network in the prediction model by extracting Social Network Analysis (SNA) features. The use of SNA enhanced the performance of the model from 84 to 93.3% against AUC standard. The model was prepared and tested through Spark environment by working on a large dataset created by transforming big raw data provided by SyriaTel telecom company. The dataset con- tained all customers’ information over 9 months, and was used to train, test, and evalu- ate the system at SyriaTel. The model experimented four algorithms: Decision Tree, Random Forest, Gradient Boosted Machine Tree “GBM” and Extreme Gradient Boosting “XGBOOST”. However, the best results were obtained by applying XGBOOST algorithm.
This algorithm was used for classification in this churn predictive model. Keywords: Customer churn prediction, Churn in telecom, Machine learning, Feature selection, Classification, Mobile Social Network Analysis, Big data
Open Access
© The Author(s) 2019. This article is distributed under the terms of the Creative Commons Attribution 4.0 International License (http://creativecommons.org/licenses/by/4.0/), which permits unrestricted use, distribution, and reproduction in any medium, provided you give appropriate credit to the original author(s) and the source, provide a link to the Creative Commons license, and indicate if changes were made.
Ahmad Et Al. J Big Data (2019) 6:28
third strategy, companies have to decrease the potential of customer’s churn, known as “the customer movement from one provider to another” . Customers’ churn is a considerable concern in service sectors with high competi- tive services. On the other hand, predicting the customers who are likely to leave the company will represent potentially large additional revenue source if it is done in the early phase .
Many research confirmed that machine learning technology is highly efficient to predict this situation. This technique is applied through learning from previous data [6, 7].
The data used in this research contains all customers’ information throughout nine months before baseline. The volume of this dataset is about 70 Terabyte on HDFS “Hadoop Distributed File System”, and has different data formats which are structured, semi-structured, and unstructured. The data also comes very fast and needs a suitable big data platform to handle it. The dataset is aggregated to extract features for each customer.
We built the social network of all the customers and calculated features like degree centrality measures, similarity values, and customer’s network connectivity for each cus- tomer. SNA features made good enhancement in AUC results and that is due to the con- tribution of these features in giving more different information about the customers.
We focused on evaluating and analyzing the performance of a set of tree-based machine learning methods and algorithms for predicting churn in telecommunications companies. We have experimented a number of algorithms such as Decision Tree, Ran- dom Forest, Gradient Boost Machine Tree and XGBoost tree to build the predictive model of customer Churn after developing our data preparation, feature engineering, and feature selection methods.
There are two telecom companies in Syria which are SyriaTel and MTN. SyriaTel com- pany was interested in this field of study because acquiring a new customer costs six times higher than the cost of retaining the customer likely to churn. The dataset pro- vided by SyriaTel had many challenges, one of them was unbalance challenge, where the churn customers’ class was very small compared to the active customers’ class. We experimented three scenarios to deal with the unbalance problem which are oversam- pling, undersampling and without re-balancing. The evaluation was performed using the Area under receiver operating characteristic curve “AUC” because it is generic and used in case of unbalanced datasets .
Many previous attempts using the Data Warehouse system to decrease the churn rate in SyriaTel were applied. The Data Warehouse aggregated some kind of telecom data like billing data, Calls/SMS/Internet, and complaints. Data Mining techniques were applied on top of the Data Warehouse system, but the model failed to give high results using this data. In contrast, the data sources that are huge in size were ignored due to the complex- ity in dealing with them. The Data Warehouse was not able to acquire, store, and process that huge amount of data at the same time. In addition, the data sources were from differ- ent types, and gathering them in Data Warehouse was a very hard process so that adding new features for Data Mining algorithms required a long time, high processing power, and more storage capacity. On the other hand, all these difficult processes in Data Ware- house are done easily using distributed processing provided by big data platform.
Ahmad Et Al. J Big Data (2019) 6:28
Furthermore, big social networks, as those in SyriaTel, are considered one of the fun- damental components of big data network graphs . The computational complexity of SNA measures is very high due to the nature of the iterative calculations done on a big scale graph, as mentioned in Eqs. (1) and (2). A lot of work to decrease the complexity of computing SNA measures has been done. For example, Barthelemy proposed a new algorithm to reduce the complexity of calculating the Betweenness centrality from O(n3) to O(n2). Elisabetta also proposed an approximation method to compute the Betweenness with less complexity. In spite of that, the traditional Data Warehouse sys- tem still suffers from deficiencies in computing the essential SNA measures on large scale networks.
Big data system allowed SyriaTel Company to collect, store, process, aggregate the data easily regardless of its volume, variety, and complexity. In addition, it enabled extracting richer and more diverse features like SNA features that provide additional information to enhance the churn predictive model.
We believe that big data facilitated the process of feature engineering which is one of the most difficult and complex processes in building predictive models. By using the big data platform, we give the power to SyriaTel company to go farther with big data sources. In addition, the company becomes able to extract the Social Network Analysis features from a big scale social graph which is built from billions of edges (transactions) that connect millions of nodes (customers). The hardware and the design of the big data platform illustrated in “Proposed churn method” section fit the need to compute these features regardless of their complexity on this big scale graph.
The model also was evaluated using a new dataset and the impact of this system to the decision to churn was tested. The model gave good results and was deployed to production.
Related Work
Many approaches were applied to predict churn in telecom companies. Most of these approaches have used machine learning and data mining. The majority of related work focused on applying only one method of data mining to extract knowledge, and the oth- ers focused on comparing several strategies to predict churn.
Gavril et al. presented an advanced methodology of data mining to predict churn for prepaid customers using dataset for call details of 3333 customers with 21 features, customer. The author applied principal component analysis algorithm “PCA” to reduce data dimensions. Three machine learning algorithms were used: Neural Networks, Sup- port Vector Machine, and Bayes Networks to predict churn factor. The author used AUC to measure the performance of the algorithms. The AUC values were 99.10%, 99.55% and 99.70% for Bayes Networks, Neural networks and support vector machine, respec- tively. The dataset used in this study is small and no missing values existed.
He et al. proposed a model for prediction based on the Neural Network algorithm in order to solve the problem of customer churn in a large Chinese telecom company which contains about 5.23 million customers. The prediction accuracy standard was the overall accuracy rate, and reached 91.1%.
Ahmad Et Al. J Big Data (2019) 6:28
Idris proposed an approach based on genetic programming with AdaBoost to model the churn problem in telecommunications. The model was tested on two standard data sets. One by Orange Telecom and the other by cell2cell, with 89% accu- racy for the cell2cell dataset and 63% for the other one.
Huang et al. studied the problem of customer churn in the big data platform. The goal of the researchers was to prove that big data greatly enhance the process of predicting the churn depending on the volume, variety, and velocity of the data. Deal- ment at China’s largest telecommunications company needed a big data platform to engineer the fractures. Random Forest algorithm was used and evaluated using AUC.
Makhtar et al. proposed a model for churn prediction using rough set theory in telecom. As mentioned in this paper Rough Set classification algorithm outper- formed the other algorithms like Linear Regression, Decision Tree, and Voted Percep- tion Neural Network.
Various researches studied the problem of unbalanced data sets where the churned customer classes are smaller than the active customer classes, as it is a major issue in churn prediction problem. Amin et al. compared six different sampling tech- niques for oversampling regarding telecom churn prediction problem. The results showed that the algorithms (MTDF and rules-generation based on genetic algo- rithms) outperformed the other compared oversampling algorithms.
Burez and Van den Poel studied the problem of unbalance datasets in churn pre- diction models and compared performance of Random Sampling, Advanced Under- Sampling, Gradient Boosting Model, and Weighted Random Forests. They used (AUC, Lift) metrics to evaluate the model. the result showed that undersampling technique outperformed the other tested techniques.
We did not find any research interested in this problem recorded in any telecommu- nication company in Syria. Most of the previous research papers did not perform the feature engineering phase or build features from raw data while they relied on ready features provided either by telecom companies or published on the internet.
In this paper, the feature engineering phase is taken into consideration to create our own features to be used in machine learning algorithms. We prepared the data using a big data platform and compared the results of four trees based machine learning algorithms.
Data Set
There are many types of data in SyriaTel used to build the churn model. These types
Are Classified As Follow:
1. Customer data It contains all data related to customer’s services and contract infor- mation. In addition to all offers, packages, and services subscribed to by the cus- tomer. Furthermore, it also contains information generated from CRM system like (all customer GSMs, Type of subscription, birthday, gender, the location of living and more ...).
Ahmad Et Al. J Big Data (2019) 6:28
2. Towers and complaints database The information of action location is represented as digits. Mapping these digits with towers’ database provides the location of this trans- action, giving the longitude and latitude, sub-area, area, city, and state.
Complaints’ database provides all complaints submitted and statistics inquiries related to coverage, problems in offers and packages, and any problem related to the telecom business.
3. Network logs data Contains the internal sessions related to internet, calls, and SMS for each transaction in Telecom operator, like the time needed to open a session for the internet and call ending status. It could indicate if the session dropped due to an error in the internal network.
4. Call details records “CDRs” Contain all charging information about calls, SMS, MMS, and internet transaction made by customers. This data source is generated as text files.
5. Mobile IMEI information It contains the brand, model, type of the mobile phone and if it’s dual or mono SIM device. This data has a large size and there is a lot of detailed information about it. We spent a lot of time to understand it and to know its sources and storing format. In addi- tion to these records, the data must be linked to the detailed data stored in relational databases that contain detailed information about the customer. The nine months of data sets contained about ten million customers. The total number of columns is about ten thousand columns.
Data exploration and challenges with SyriaTel dataset Spark engine is used to explore the structure of this dataset, it was necessary to make the exploration phase and make the necessary pre-preparation so that the dataset becomes suitable for classification algorithms. After exploring the data, we found that about 50% of all numeric variables contain one or two discrete values, and nearly 80% of all the categorical variables have Less than 10 categories, 15% of the numeri- cal variables and 33% of the categorical variables have only one value. Most of some variables’ values are around zero. We found that 77% of the numerical variables have more than 97% of their values filled with 0 or null value. These results indicate that a large number of variables can be removed because these variables are fixed or close to a constant. This dataset encounters many challenges as follow.
Data Volume
Since we don’t know the features that could be useful to predict the churn, we had to work on all the data that reflect the customer behavior in general. We used data sets related to calls, SMS, MMS, and the internet with all related information like com- plaints, network data, IMEI, charging, and other. The data contained transactions for all customers during nine months before the prediction baseline. The size of this data was more than 70 Terabyte, and we couldn’t perform the needed feature engineering phase using traditional databases.
Data Variety
The data used in this research is collected from multiple systems and databases. Each source generates the data in a different type of files as structured, semi-structured (XML-JSON) or unstructured (CSV-Text). Dealing with these kinds of data types is very hard without big data platform since we can work on all the previous data types with- out making any modification or transformation. By using the big data platform, we no longer have any problem with the size of these data or the format in which the data are represented.
Unbalanced Dataset
The generated dataset was unbalanced since it is a special case of the classification prob- lem where the distribution of a class is not usually homogeneous with other classes. The dominant class is called the basic class, and the other is called the secondary class. The data set is unbalanced if one of its categories is 10% or less compared to the other one .
Although machine learning algorithms are usually designed to improve accuracy by reducing error, not all of them take into account the class balance, and that may give bad results . In general, classes are considered to be balanced in order to be given the same importance in training.
We found that SyriaTel dataset was unbalanced since the percentage of the secondary class that represents churn customers is about 5% of the whole dataset.
Extensive Features
The collected data was full of columns, since there is a column for each service, prod- uct, and offer related to calls, SMS, MMS, and internet, in addition to columns related to personnel and demographic information. If we need to use all these data sources the number of columns for each customer before the data being processed will exceed ten thousand columns.
Missing Values
There is a representation of each service and product for each customer. Missing values may occur because not all customers have the same subscription. Some of them may have a number of services and others may have something different. In addition, there are some columns related to system configurations and these columns have only null value for all customers.
Proposed Churn Method
In order to build the churn predictive system at SyriaTl, a big data platform must be installed. Hortonworks Data Platform (HDP)1 was chosen because it is a free and an open source framework. In addition, it is under the Apache 2.0 License. HDP plat- form has a variety of open source systems and tools related to big data. These open source systems and tools are integrated with each other. Figure 1 presents the ecosystem 1 https://hortonworks.com/.
Ahmad Et Al. J Big Data (2019) 6:28
of HDP, where each group of tools is categorized under specific specialization like Data Management, Data Access, Security, Operations and Governance Integration. The installation of HDP framework was customized in order to have the only needed tools and systems that are enough to go through all phases of this work. This custom- ized package of installed systems and tools is called SYTL-BD framework (SyriaTel’s big data framework). We installed Hadoop Distributed File System HDFS2 to store the data, Spark execution engine3 to process the data, Yarn4 to manage the resources, Zeppelin5 as the development user interface, Ambari6 to monitor the system, Ranger7 to secure the system and (Flume8 System and Scoop9 tool) to acquire the data from outside SYTL-BD framework into HDFS.
The used hardware resources contained 12 nodes with 32 Gigabyte RAM, 10 Terabyte storage capacity, and 16 cores processor for each node. A nine consecutive months data- set was collected. This dataset will be used to extract the features of churn predictive model. The data life cycle went through several stages as shown in Fig. 2 Spark engine was used in most of the phases of the model like data processing, feature engineering, training and testing the model since it performs the processing on RAM. In addition, there are many other advantages. One of these advantages is that this engine containing a variety of libraries for implementing all stages of machine learning lifecycle.
Data Acquisition And Storing
Moving the data from outside SYTL-BD into HDFS was the first step of work. The data is divided into three main types which are structured, semi-structured and unstructured. Apache Flume is a distributed system used to collect and move the unstructured (CSV and text) and semi-structured (JSON and XML) data files to HDFS. Figure 3 shows Fig. 1 Hortonworks data platform HDP—big data framework 2 https://hadoop.apache.org/docs/r1.2.1/hdfs_design.html.
3 https://spark.apache.org/. 4 https://hadoop.apache.org/docs/current/hadoop-yarn/hadoop-yarn-site/YARN.html. 5 https://zeppelin.apache.org/.
6 https://ambari.apache.org/. 7 https://ranger.apache.org/. 8 https://flume.apache.org.
9 https://sqoop.apache.org/.
Ahmad Et Al. J Big Data (2019) 6:28
the designed architecture of flume in SYTL-BD. There are three main components in FLUME. These components are the data Source, the Channel where the data moves and the Sink where the data is transported.
Flume agents transporting files exist in the defined Spooling Directory Source using one channel, as configured in SYTL-BD. This channel is defined as Memory Channel because it performed better than the other channels in FLUME. The data moves across the channel to be finally written in the sink which is HDFS. The data transformed to HDFS keep in the same format type as it was.
Apache SQOOP is the distributed tool used to transfer the bulk of data between HDFS and relational databases (Structured data). This tool was used to transfer all the data which exists in databases into HDFS by using Map jobs. Figure 4 shows the archi- tecture of SQOOP import process where four mappers are defined by default. Each Map job selects part of the data and moves it to HDFS. The data is saved in CSV file type after being transported by SQOOP to HDFS.
After transporting all the data from its sources into HDFS, it was important to choose the appropriate file type that gives the best performance in regards to space utilization and execution time. This experiment was done using spark engine where Fig. 2 Proposed churn Prediction System Architecture Fig. 3 Apache Flume configured system architecture
Ahmad Et Al. J Big Data (2019) 6:28
Data Frame library10 was used to transform 1 terra byte of CSV data into Apache Par- quet11 file type and Apache Avro12 file type. In addition to that, three compression scenarios were taken into consideration in this experiment.
Parquet file type was the chosen format type that gave the best results. It is a colum- nar storage format since it has efficient performance compared with the others, espe- cially in dealing with feature engineering and data exploration tasks. On the other hand, using Parquet file type with Snappy Compression technique gave the best space utilization. Figure 5 shows some comparison between file types.
Feature Engineering
The data was processed to convert it from its raw status into features to be used in machine learning algorithms. This process took the longest time due to the huge numbers of columns. The first idea was to aggregate values of columns per month (average, count, sum, max, min ...) for each numerical column per customer, and the count of distinct values for categorical columns.
Another type of features was calculated based on the social activities of the custom- ers through SMS and calls. Spark engine is used for both statistical and social fea- tures, the library used for SNA features is the Graph Frame.
Fig. 4 Apache Sqoop Data Import Architecture
10 https://spark.apache.org/docs/latest/sql-programming-guide.html. 11 https://parquet.apache.org/. 12 https://avro.apache.org/.
Ahmad Et Al. J Big Data (2019) 6:28
• Statistics features These features are generated from all types of CDRs, such as the average of calls made by the customer per month, the average of upload/down- load internet access, the number of subscribed packages, the percentage of Radio Access Type per site in month, the ratio of calls count on SMS count and many features generated from aggregating data of the CDRs.
Since we have data related to all customers’ actions in the network, we aggre- gated the data related to Calls, SMS, MMS, and internet usage for each customer per day, week, and month for each action during the nine months. Therefore, the number of generated features increased more than three times the number of the columns. In addition, we entered the features related to complaints submitted from the customers from all systems. Some features were related to the number of complaints, the percentage of coverage complaints to the whole complaints sub- mitted, the average duration between each two complaints sequentially, the dura- tion in “Hours” to close the complaint, the closure result, and other features.
The features related to IMEI data such as the type of device, the brand, dual or mono device, and how many devices the customer changed were extracted. We did many rounds of brainstorming with seniors in the marketing section to decide what features to create in addition to those mentioned in some researches.
We created many features like percentage of incoming/out-coming calls, SMS, MMS to the competitors and landlines, binary features to show if customers were subscribing some services or not, rate of internet usage between 2G, 3G and 4G, number of devices used each month, number of days being out of coverage, per- centage of friends related to competitor, and hundred of other features.
Figures 6 and 7 visualize some of the basic categorical and numerical features to give more insight on the deference between churn and non-churn classes. Fig. 5 Differences in space utilization and execution time per file type
Ahmad Et Al. J Big Data (2019) 6:28
• Social Network Analysis features Data transformation and preparation are per- formed to summarize the connections between every two customers and build a social network graph based on CDR data taken for the last 4 months. Graph frame library on spark is used to accomplish this work. The social network graph con- sists of Nodes and edges.
Fig. 6 Distribution of some main categorical features Fig. 7 Feature distribution for some main numerical features. Panel (a) visualizes the distribution of Day of Last Outgoing Transaction feature. Panel (b) visualizes the feature distribution of Average Radio Access Type Between 3G and 2G. Panel (c) also visualizes the distribution of Total Balance feature. Panel (d) shows the feature distribution of Percentage Transaction with other operators. Similarly, panel (e) visualizes the distribution of Percentage of Signaling Error/Dropped calls. Finally, panel (f) visualizes the distribution of the GSM Age feature. The red color is used in all panels to represent the churned customers’ class and the blue
Ahmad Et Al. J Big Data (2019) 6:28
• Nodes: represent GSM number of subscribers. • Edges: represent interactions between subscribers (Calls, SMS, and MMS). The graph edges are directed since we have A to B and B to A.
Figure 8 visualizes a sample of the build social network in SyriaTel where the red nodes are SyriaTel’s customers and the Yellow nodes are MTN’s Customers, the lines between the nodes express the interaction between the nodes.
The total social graph contained about 15 million nodes that represent SyriaTel, MTN, and Baseline numbers and more than 2.5 Billion edges. Graph-based features are extracted from the social graph. The graph is a weighted directed graph. We built three graphs depending on the used edges’ weight. The weight of edges is the number of shared events between every two customers. We used three types of weights: (1) the normalized calling duration between customers, (2) the normal- ized total number of calls, SMS, and MMS, (3) the mean of the previous two normalized weights. The normalization process varies according to the algorithm used to extract the features as we see in the formulas of these algorithms. Based on the directed graphs, we use PageRank , Sender Rank algorithms to produce two features for each graph.
• The weighted Page Rank equation is defined as follows
N′∈N(N) Wn→N′ Pr(N)
Fig. 8 Visualization for a sample of the Syrian social community
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.