Visualization
Data visualization is the use of visual representations to enhance our understanding of data, facilitate hypothesis building and decision making, and, more importantly, communicate our analysis findings. Therefore, data visualization is an essential skill for a data scientist. This post focuses on accurately summarizing data and reporting findings using Matplotlib and Pandas plots.
In the context of machine learning, we use visualization data for three main reasons:
- Exploratory Analysis: The objective of exploratory data visualization is to get ready for analysis and to gain a better understanding of the data. We use histograms, box plots to investigate the distribution of the data and identify outlier or unusual cases, and scatter plots to investigate relationships between variables. The exploratory visual analysis helps us to identify important variables in our dataset, data points to exclude, and types of analyses to perform.
- Verification & Validation: A machine learning algorithm will learn the data and provide an answer. But how reliable is the answer? We usually train and test our algorithms using different parameters and rely on visual aids to validate the output. These visual aids can tell us a better story than a simple mean squared error. For example, we can observe in which part of the data our model fails to predict.
- Communications: We usually rely on charts and demonstrations to communicate our findings. This type of visualization is usually for non-experts. Therefore, we need to find creative ways to make complex data and findings more understandable to the average user in a few charts and diagrams.
Matplotlib Basics
Matplotlib is a comprehensive library for creating visualizations in Python. Matplotlib is also the basis for other visualization libraries and provides many types of plotting functions. All Matplotlib functions take a numpy array as the input, and some of them also accept Pandas dataframes. Unfortunately, Matplotlib is not a one-for-all solution. We may need to write additional coding to create plots suitable for data analysis. Therefore, we will rely on Pandas and Seaborn libraries to create plots for analysis.
In the following code, we plot y=10x and y=x2 functions for x=0,0.05,0.1,1.5,…,10 using the plot() function. Given two arrays x and y of the same length, plot(x,y) plots y values versus x values as lines and/or markers. Array x is also optional, and if x is omitted, plot(y) will plot y values sequentially.
We will always start creating the plot with a figure() statement. The figure() creates a top-level container that holds all elements required for the plot. When a figure() statement is executed, it creates a current figure. Any additions or changes that we make after a figure() statement apply to the current figure until the next figure() statement is run.
The show() function displays all figures. The show() statement will not have any effect if you run the code in Spyder. However, we will use the show function in case we want to run your code in a Jupyter Notebook.
import numpy as np
import matplotlib.pyplot as plt
#create some data to plot
x=np.linspace(0,10,21,endpoint=True)
y1=10*x
y2=x**2
plt.figure()
#plot y1 array with a blue (b) blue line
plt.plot(x, y1, c='b', label='a linear function')
#plot y2 array with red (r) redline
plt.plot(x, y2, c='r', label='a quadratic function')
plt.ylabel('y') #set the label of the y axis
plt.xlabel('x') #set the label of the x axis
plt.title('Function Plots')
plt.grid(True) #show grids
plt.xlim(0, 10) #minimum and maximum of the x axis
plt.ylim(0, 110) #minimum and maximum of the y axis
plt.xticks(range(0,11)) #set the x axis tick marks
plt.yticks(range(0,110,10)) #set the y axis tick marks
plt.legend() # show the legend
plt.show() #show the figure
Figures, Axes, Axis
The region that displays the data is called Axes. In fact, we create plots on Axes, not on a Figure. The Figure is a just container that holds all the child Axes. There are essentially two ways to create Matplotlib plots.
- Use pyplot functions for plotting (the approach used above).
- Explicitly create figures and axes, and call methods on them (the “object-oriented (OO) style”).
The example below illustrates how we can create the same plot using the OO style.
fig, ax=plt.subplots() # Create a figure and an axes.
ax.plot(x, y1, label='a linear function') # Plot some data on the axes.
ax.plot(x, y2, label='a quadratic function') # Plot more data on the axes.
ax.set_xlabel('$x$') # Add an x-label to the axes.
ax.set_ylabel('$y$') # Add a y-label to the axes.
ax.set_title("Function Plots") # Add a title to the axes.
ax.set_xticks(range(0,11)) #set the x axis tick marks
ax.set_yticks(range(0,110,10)) #set the y axis tick marks
ax.grid(True)
ax.legend() # Add a legend.
All attributes and properties of the axes are given in the Matplotlib documentation: https://matplotlib.org/stable/api/axes_api.html#axis-limits
The main advantage of using Axes is to plot different data on different Axes using different properties and styles. This provides the flexibility of organizing the chart in different layouts. See the example below.
fig, ax=plt.subplots(1,2) # Create a figure and an axes.
ax[0].plot(x, y1, label='a linear function') # Plot some data on the axes.
ax[1].plot(x, y2, label='a quadratic function') # Plot more data on the axes.
ax[0].set_xlabel('$x$') # Add an x-label to the axes.
ax[0].set_ylabel('$10x$') # Add a y-label to the axes.
ax[1].set_xlabel('$x$') # Add an x-label to the axes.
ax[1].set_ylabel('$x^2$') # Add a y-label to the axes.
ax[0].set_title("Function Plots") # Add a title to the axes.
ax[0].set_xticks(range(0,11)) #set the x axis tick marks
ax[0].set_yticks(range(0,110,10)) #set the y axis tick marks
ax[0].grid(True)
ax[0].legend() # Add a legend.