Chris Parmer — home

Cohort retention heatmap

UCI Online Retail transactions

Example from the compendium of canonical charts

Cohort retention heatmap — UCI Online Retail transactions

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