Mathew K Analytics

Lesson 18 · Data visualisation in python

Understanding Heatmaps and Correlation Matrices with Seaborn for Data Analysis

Let us start exploring how to use Python and Seaborn to build heatmaps and visualize relationships in data. By the end, you will be able to: Load a real…

⬇ Download notebookOpen in Colab ↗

What you'll learn

Data

No separate download needed — the notebook creates or downloads everything it uses.

📓 Full notebook

Download .ipynb

Heatmaps & Correlation Matrices in Python: Beginner Workshop#

Let us start exploring how to use Python and Seaborn to build heatmaps and visualize relationships in data.

By the end, you will be able to:

  • Load a real dataset
  • Draw a basic heatmap
  • Understand correlation matrices
  • Solve a fun data mystery!
# Import required libraries and suppress warnings
import warnings
from statsmodels.tools.sm_exceptions import ConvergenceWarning, ValueWarning
warnings.filterwarnings("ignore", category=ValueWarning)
warnings.filterwarnings("ignore", category=ConvergenceWarning)
warnings.filterwarnings("ignore", category=RuntimeWarning)  # e.g., from log/exp on edge values
import warnings; warnings.filterwarnings('ignore')
import pandas as pd
import seaborn as sns
import matplotlib.pyplot as plt

Step 1: Loading a Real Dataset#

We will use the Airline Passengers dataset. It tracks the number of airline travelers each month.

Let us load it and peek at the numbers.

# Data setup: download and load with pandas
url = 'https://raw.githubusercontent.com/jbrownlee/Datasets/master/airline-passengers.csv'
df = pd.read_csv(url)

# Show shape and first 5 rows
print('Shape of table:', df.shape)
df.head()
Shape of table: (144, 2)
Month Passengers
0 1949-01 112
1 1949-02 118
2 1949-03 132
3 1949-04 129
4 1949-05 121
# Quick line plot to visualize time series data
plt.figure(figsize=(8,4))
plt.plot(df['Month'], df['Passengers'], marker='o')
plt.title('Passengers Over Time')
plt.xlabel('Month')
plt.ylabel('Number of Passengers')
plt.xticks(rotation=45)
plt.tight_layout()
plt.show()
No description has been provided for this image

Step 2: What is a Heatmap?#

A heatmap is a grid where colors show the differences in numbers or relationships. Darker or lighter colors mean bigger or smaller values.

They are used to spot patterns, trends, or high and low spots in data.

Let us make one!

# Make up a tiny grid of numbers to illustrate heatmap basics
import numpy as np
data = np.array([[1, 4, 3],
                 [2, 0, 8],
                 [6, 3, 5]])
sns.heatmap(data, annot=True, cmap='YlGnBu')
plt.title('Toy Example: Heatmap of Numbers')
plt.show()
No description has been provided for this image

Step 3: The Correlation Matrix#

A correlation matrix compares many columns to see how closely related they are.

Each cell shows a number called 'correlation'. A value of +1 means a perfect positive link. -1 means a perfect negative link.

In the real world, this is how we discover which variables rise and fall together.

# Create a small DataFrame with random numbers for demo
np.random.seed(0)
mini_df = pd.DataFrame({
    'A': np.random.rand(8),
    'B': np.random.rand(8),
    'C': np.random.rand(8)*2 + 1
})
mini_corr = mini_df.corr()
print('Small correlation matrix:')
print(mini_corr.round(2))
Small correlation matrix:
      A     B     C
A  1.00 -0.22  0.18
B -0.22  1.00 -0.23
C  0.18 -0.23  1.00
# Draw a heatmap of the correlation matrix
sns.heatmap(mini_corr, annot=True, cmap='coolwarm', center=0)
plt.title('Correlation Matrix Heatmap (Mini Example)')
plt.show()
No description has been provided for this image

Step 4: Heatmaps for Real Data#

Let us use our airline passenger data now!

First, we will add a new column: which year is it? Then we will see the average number of passengers for each year and month.

# Extract year and month columns
df['Year'] = df['Month'].str[:4]
df['MonthNum'] = df['Month'].str[5:7]

# Group by year and month, get mean passenger count
pivot = df.pivot_table(values='Passengers', index='Year', columns='MonthNum', aggfunc='mean')
pivot.head()
MonthNum 01 02 03 04 05 06 07 08 09 10 11 12
Year
1949 112.0 118.0 132.0 129.0 121.0 135.0 148.0 148.0 136.0 119.0 104.0 118.0
1950 115.0 126.0 141.0 135.0 125.0 149.0 170.0 170.0 158.0 133.0 114.0 140.0
1951 145.0 150.0 178.0 163.0 172.0 178.0 199.0 199.0 184.0 162.0 146.0 166.0
1952 171.0 180.0 193.0 181.0 183.0 218.0 230.0 242.0 209.0 191.0 172.0 194.0
1953 196.0 196.0 236.0 235.0 229.0 243.0 264.0 272.0 237.0 211.0 180.0 201.0
# Heatmap of passenger numbers by month and year
plt.figure(figsize=(10,6))
sns.heatmap(pivot, annot=True, fmt='.0f', cmap='YlOrRd')
plt.title('Passengers by Year and Month: Heatmap')
plt.xlabel('Month')
plt.ylabel('Year')
plt.show()
No description has been provided for this image

Step 5: Correlation Matrix for the Whole Dataset#

Now, let us check how columns in the whole dataset relate. This will be simple since the dataset is small, but it is a good habit!

# Calculate correlation matrix for all numeric columns
# Ensure only numeric columns are included for correlation
numeric_cols = df.select_dtypes(include='number')
corr = numeric_cols.corr()
print('Correlation matrix for passengers data:')
print(corr.round(2))
Correlation matrix for passengers data:
            Passengers
Passengers         1.0
# Heatmap for the dataset's correlation matrix
sns.heatmap(corr, annot=True, cmap='vlag', center=0)
plt.title('Passengers Dataset Correlation Heatmap')
plt.show()
No description has been provided for this image

Step 6: Mini-Project - Heatmap Data Detective#

Suppose you work for an airline and want to spot your quietest and busiest times fast.

Use a heatmap to find:

  • The month with the highest average passengers
  • The month with the lowest

Let us find clues together!

# Get month averages across all years
month_means = df.groupby('MonthNum')['Passengers'].mean()
print('Average passengers by month:')
print(month_means.sort_values(ascending=False))
Average passengers by month:
MonthNum
07    351.333333
08    351.083333
06    311.666667
09    302.416667
05    271.833333
03    270.166667
04    267.083333
10    266.583333
12    261.833333
01    241.750000
02    235.000000
11    232.833333
Name: Passengers, dtype: float64
# Visualize with a bar heatmap
plt.figure(figsize=(8,2))
sns.heatmap([month_means.values], annot=True, cmap='YlOrBr', cbar=True)
plt.yticks([])
plt.xticks(ticks=np.arange(12)+.5, labels=month_means.index, rotation=0)
plt.title('Average Passengers by Month: Heatmap')
plt.xlabel('Month (MM)')
plt.show()
No description has been provided for this image

Step 7: Troubleshooting Common Heatmap Issues#

If you see errors or blank charts:

  • Check for missing data (NaN)
  • Make sure your numbers are not all the same
  • Look at color and labels

The more you practice, the clearer these will be!

# Check for missing values
print('Has missing data?')
print(df.isnull().any())
Has missing data?
Month         False
Passengers    False
Year          False
MonthNum      False
dtype: bool

Recap: Key Ideas about Heatmaps#

Heatmaps are a simple but powerful way to see patterns in data.

  • Try to add more columns to your tables
  • Practice grouping and reshaping with pandas
  • Use colors to catch your eye on the big picture

The more you experiment, the more insights youll discover!

Thanks for Learning! Your Turn Now#

Try these:

  • Build your own heatmap with other datasets
  • Change the color palette styles
  • Share your best charts in the comments below

Subscribe for more beginner Python and data tips!

Found this useful?

All lessons, notebooks and datasets here are free. If they helped you, a coffee keeps new lessons coming.