In this issue, we will introduce mainly the Matplotlib library for plotting, along with its good friend Seaborn as planned. While compiling this issue, I improved the layout system of the entire website's posts, hope you like it. I will also change some styles in this issue.
Previously, we have looked at Numpy and Pandas, and we can basically obtain and process the data. The next direct question is: what do these data actually look like? At this point, just looking at a bunch of numbers can really be a bit dizzying, so we need to plot.
In fact, Matplotlib is still very important for ML. Plotting is not just for aesthetics; many times, whether a model is overfitting, whether the data has strange distributions, or whether two variables are related can be much clearer at a glance in a graph than in a table.
Of course, Matplotlib can plot a lot of things. In this issue, we will first discuss the most commonly used items, and then we can gradually fill in the gaps when writing our own projects.
1. First, prepare Matplotlib
If you don't have matplotlib in your environment, you can install it in the terminal or Colab:
pip install matplotlib seaborn
Then we generally import it like this. Here, pyplot is basically our most commonly used plotting tool, and everyone calls it plt, so don't be nervous when you see plt later.
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import seaborn as sns
2. The most basic line graph
Let's start with the simplest line graph. Suppose we have a sin curve. Here, linspace will help us take 200 numbers between -10 and 10, and then we can just put it into np.sin() to get started.
x = np.linspace(-10, 10, 200)
y = np.sin(x)
plt.figure(figsize=(9, 5))
plt.plot(x, y, color="royalblue", linewidth=2, label="sin(x)")
plt.axhline(0, color="gray", linewidth=0.8)
plt.title("A simple sine curve")
plt.xlabel("x")
plt.ylabel("sin(x)")
plt.grid(alpha=0.25)
plt.legend()
plt.show()
After running it, you will see a very standard curve.figure is the entire canvas,plot is to draw lines,xlabel and ylabel are the names of the axes, and the final show() is to actually display the graph. In fact, none of these lines are strictly necessary, but adding title, label, and grid makes it much clearer. A graph that you can understand is also easier for others to understand.
3. Scatter plot: to see if there is any relationship between two things
I think scatter plots are very common in ML. For example, we want to see if there is any relationship between study hours and score. To make the example runnable directly, I will use random data to simulate it.
rng = np.random.default_rng(7)
study_hours = rng.uniform(1, 10, 90)
score = 42 + study_hours * 5.4 + rng.normal(0, 7, 90)
plt.figure(figsize=(8, 5))
plt.scatter(study_hours, score, color="#ff7a59", alpha=0.72)
plt.title("Study hours and score")
plt.xlabel("Study hours")
plt.ylabel("Score")
plt.grid(alpha=0.2)
plt.show()
In this graph, each dot represents a piece of data.alpha is the transparency; if there are many dots, adding a bit of transparency will make it easier to see where it is particularly dense.
This example will likely show an upward trend, but that doesn't mean studying longer will definitely lead to a higher score; it just appears that there is a correlation in this dataset. It's important to distinguish this; correlation does not imply causation, and caution is still needed when doing data analysis later.
4. Bar chart and histogram
Bar charts are suitable for comparing different categories, such as the accuracy of different models. Histograms, on the other hand, show the distribution of a number. These two types of graphs may seem basic, but they are really used in model evaluation.
models = ["Logistic Regression", "Decision Tree", "Random Forest"]
accuracy = [0.79, 0.83, 0.89]
plt.figure(figsize=(9, 5))
plt.bar(models, accuracy, color=["#90caf9", "#81c784", "#ffcc80"])
plt.ylim(0, 1)
plt.title("Model accuracy comparison")
plt.ylabel("Accuracy")
plt.show()
values = rng.normal(loc=70, scale=12, size=500)
plt.figure(figsize=(8, 5))
plt.hist(values, bins=24, color="#8e7dff", edgecolor="white")
plt.title("A simple distribution")
plt.xlabel("Value")
plt.ylabel("Count")
plt.show()
Note ylim(0, 1) is just because accuracy is inherently between 0 and 1. If the ranges of different charts are not the same, one should not arbitrarily cut the y-axis to make one model look impressive; this is something to be mindful of when making reports.
5. Putting several graphs in one figure
Sometimes we don't want to open many images, we can use subplots. The example below puts a line graph and a histogram in the same figure. When I first used this, I also had a bit of trouble distinguishing fig And axcan be simply understood as fig is the entire sheet of paper,ax is each small area on the paper.
fig, ax = plt.subplots(1, 2, figsize=(12, 4))
ax[0].plot(x, y, color="#0071e3")
ax[0].set_title("sin(x)")
ax[0].set_xlabel("x")
ax[1].hist(values, bins=22, color="#34c759", edgecolor="white")
ax[1].set_title("Value distribution")
ax[1].set_xlabel("Value")
plt.tight_layout()
plt.show()
tight_layout() This thing is very convenient; it will automatically adjust the spacing for us, otherwise the labels sometimes collide. It's not a very deep function, but it's very practical.
6. Seaborn: A layer on top of Matplotlib
Seaborn is actually built on top of Matplotlib, but its default style looks a bit nicer and it works very well with Pandas DataFrame. For example, with the small DataFrame below, we can directly hand it over to Seaborn.
scores = pd.DataFrame({
"study_hours": [2, 3, 4, 4, 5, 6, 7, 8, 9, 10],
"score": [52, 55, 63, 60, 68, 74, 73, 82, 88, 91],
"group": ["A", "A", "B", "A", "B", "B", "A", "B", "A", "B"],
})
sns.set_theme(style="whitegrid")
plt.figure(figsize=(8, 5))
sns.scatterplot(data=scores, x="study_hours", y="score", hue="group", s=90)
plt.title("Scores in two groups")
plt.show()
Here hue="group" will automatically change different groups into different colors. Seaborn has this feeling: for many common data visualizations, it has already set up some configurations for us. Of course, this doesn't mean Matplotlib is useless; Seaborn ultimately returns to Matplotlib, so using both together is more comfortable.
Another example that I think will be used easily in the future is the heatmap. For instance, let's first look at the correlation between several variables:
data = pd.DataFrame({
"hours": [2, 3, 4, 5, 6, 7, 8, 9],
"practice": [1, 2, 2, 3, 4, 5, 6, 7],
"score": [50, 55, 59, 67, 72, 78, 84, 90],
})
plt.figure(figsize=(6, 4))
sns.heatmap(data.corr(numeric_only=True), annot=True, cmap="coolwarm", vmin=-1, vmax=1)
plt.title("Correlation heatmap")
plt.show()
annot=True will write the numbers directly in the boxes,cmap is the color. Red-blue color maps are common for visualizing correlation, but it should be reiterated that correlation is just a starting point, not the final answer.
7. For ML, the most worthwhile thing to plot
I think the most worthwhile thing to plot right at the beginning of training a model is the training loss and validation loss. Because this chart can roughly tell us whether the model is still learning or if it has started to overfit. The data below is just simulated; during actual training, just append the loss of each epoch.
epochs = np.arange(1, 11)
train_loss = [0.92, 0.75, 0.61, 0.49, 0.42, 0.36, 0.31, 0.28, 0.25, 0.22]
val_loss = [0.95, 0.79, 0.66, 0.56, 0.51, 0.49, 0.50, 0.54, 0.60, 0.69]
plt.figure(figsize=(9, 5))
plt.plot(epochs, train_loss, marker="o", label="training loss")
plt.plot(epochs, val_loss, marker="o", label="validation loss")
plt.xlabel("Epoch")
plt.ylabel("Loss")
plt.title("Training and validation loss")
plt.legend()
plt.grid(alpha=0.25)
plt.show()
If the training loss keeps decreasing, but the validation loss starts to rise later, that could be a signal of overfitting. At this point, there's no need to panic immediately; you can try early stopping, dropout, data augmentation, or simply check the data first.
Many ML problems are not issues with the model itself; after plotting the graph, you may find that the data issues are even bigger, haha.
So that's it for this episode. Matplotlib and Seaborn really have a lot more to offer, such as boxplot, pairplot, confusion matrix, 3D plot, animation, etc. I won't cover everything here, otherwise this episode would turn into a super long documentation.
Later, when we do model evaluation or Deep Learning, we will continue to use these charts, and we will discuss them again then.
In the next issue, we will continue with the Model Building section and gradually move towards Deep Learning and Computer Vision.
Everyone can first take a CSV file and try to draw a line graph, scatter plot, and heatmap. It's okay if you mess it up; running it a few more times will help you get familiar with it. See you all next time~