import pandas as pd
import matplotlib.pyplot as plt
import numpy as np
df = pd.read_csv("gapminder.csv").dropna(subset=["population_millions", "life_expectancy", "continent"])
# log x 之前先过滤非正数(保险)
df = df[df["population_millions"] > 0].copy()
# 用分位数裁掉极端大国,避免把所有点压在左边(更好看)
x = df["population_millions"].astype(float)
y = df["life_expectancy"].astype(float)
x_max = np.nanpercentile(x, 99.5) # 只裁掉最极端的 0.5%
mask = x <= x_max
df = df[mask].copy()
continents = sorted(df["continent"].unique())
cmap = plt.get_cmap("tab10")
color_map = {c: cmap(i % 10) for i, c in enumerate(continents)}
plt.figure()
for c in continents:
d = df[df["continent"] == c]
plt.scatter(
d["population_millions"],
d["life_expectancy"],
s=16,
alpha=0.35,
label=c,
c=[color_map[c]],
edgecolors="none",
)
plt.xscale("log")
plt.xlabel("Population (millions, log scale)")
plt.ylabel("Life Expectancy")
plt.title("Life Expectancy vs Population (colored by continent)")
plt.legend(title="Continent", frameon=False)
plt.tight_layout()
plt.show()