Cohort retention heatmap
UCI Online Retail transactions
Example from the compendium of canonical charts
Python Code
"""Cohort retention heatmap — UCI Online Retail transactions."""
import io
import numpy as np
import pandas as pd
import plotly.graph_objects as go
import requests, warnings
URL = "https://raw.githubusercontent.com/amir-hojjati/Data-Analysis-Online-Retail-Transactions/master/Original-Dataset/Online%20Retail.csv"
def generate():
print("fetching Online Retail dataset …")
warnings.filterwarnings("ignore")
requests.packages.urllib3.disable_warnings()
# Stream first ~5MB to keep it manageable (full file ~44MB)
r = requests.get(URL, verify=False, timeout=60, stream=True)
r.raise_for_status()
buf = b""
for chunk in r.iter_content(65536):
buf += chunk
if len(buf) >= 25_000_000:
break
r.close()
text = buf.decode("latin-1", errors="ignore")
text = text[:text.rfind("\n")]
df = pd.read_csv(io.StringIO(text), encoding="latin-1")
df.columns = [c.strip() for c in df.columns]
print(f" raw cols: {list(df.columns)}, rows: {len(df)}")
# Drop credit notes and missing CustomerID
cust_col = next((c for c in df.columns if "customer" in c.lower()), None)
inv_col = next((c for c in df.columns if "invoice" in c.lower() and "date" not in c.lower()), None)
date_col = next((c for c in df.columns if "date" in c.lower()), None)
if not all([cust_col, inv_col, date_col]):
raise ValueError(f"Missing columns. Found: {list(df.columns)}")
df = df.dropna(subset=[cust_col])
df = df[~df[inv_col].astype(str).str.startswith("C")]
df[date_col] = pd.to_datetime(df[date_col], errors="coerce", dayfirst=False)
df = df.dropna(subset=[date_col])
df["CohortMonth"] = df.groupby(cust_col)[date_col].transform("min").dt.to_period("M")
df["InvoiceMonth"] = df[date_col].dt.to_period("M")
df["MonthsSince"] = (df["InvoiceMonth"] - df["CohortMonth"]).apply(lambda x: x.n)
cohort_sizes = df.groupby("CohortMonth")[cust_col].nunique()
retention = df.groupby(["CohortMonth", "MonthsSince"])[cust_col].nunique().reset_index()
retention.columns = ["CohortMonth", "MonthsSince", "ActiveCustomers"]
retention["RetentionRate"] = retention.apply(
lambda r: r["ActiveCustomers"] / cohort_sizes[r["CohortMonth"]] * 100, axis=1
)
# Pivot to matrix
pivot = retention.pivot(index="CohortMonth", columns="MonthsSince", values="RetentionRate")
pivot = pivot.sort_index()
# Limit to first 12 months and reasonable number of cohorts
pivot = pivot.iloc[:, :13]
z = pivot.values.tolist()
x = [str(c) for c in pivot.columns]
y = [str(c) for c in pivot.index]
text_vals = [[f"{v:.0f}%" if not np.isnan(v) else "" for v in row] for row in pivot.values]
fig = go.Figure(go.Heatmap(
z=z, x=x, y=y,
text=text_vals,
texttemplate="%{text}",
textfont=dict(size=16),
colorscale=COLORSCALE,
zmin=0, zmax=100,
colorbar=dict(title="Retention %", thickness=12, ticksuffix="%"),
hovertemplate="Cohort %{y}<br>Month %{x}<br>%{z:.1f}%<extra></extra>",
))
fig.update_layout(
title=dict(text="Monthly Cohort Retention — Online Retail (UCI)", x=0.5),
xaxis=dict(title="Months since first purchase", showgrid=False),
yaxis=dict(title="Cohort (first purchase month)", autorange="reversed", showgrid=False),
margin=dict(t=60, b=60, l=120, r=80),
height=max(380, len(y) * 26 + 120),
)
return fig
if __name__ == "__main__":
generate()
Made with Plotly