19. Polars#
除了 Anaconda 中已有的库之外,本讲座还需要以下库:
!pip install --upgrade polars yfinance
19.1. 概述#
Polars 是一个用 Rust 编写的快速 Python 数据处理库。
由于其性能优势,它作为 pandas 的现代替代品已经获得了广泛的关注。
Polars 在设计时充分考虑了性能和内存效率,主要利用了以下技术:
Apache Arrow 列式存储格式,实现快速的数据访问
惰性求值(lazy evaluation),用于优化查询执行
并行处理,充分利用所有可用的 CPU 核心
围绕列表达式构建的富有表现力的 API
Tip
为什么要考虑用 Polars 替代 pandas?
内存:pandas 通常需要相当于数据集大小 5–10 倍的内存;Polars 只需 2–4 倍
速度:对于许多常见操作,Polars 的速度快 10–100 倍
参见:Polars TPC-H 基准测试,可查看最新的性能对比结果
在整个讲座中,我们假设已经执行了以下导入语句
import polars as pl
import numpy as np
import matplotlib.pyplot as plt
与 Pandas 类似,Polars 定义了两种重要的数据类型:Series 和 DataFrame。
你可以将 Series 理解为一列数据,例如某个变量的一组观测值。
DataFrame 是一个二维对象,用于存储若干相关联的数据列。
19.2. Series#
我们先从 Series 开始。
首先创建一个由四个随机观测值组成的 series
s = pl.Series(name='daily returns', values=np.random.randn(4))
s
| daily returns |
|---|
| f64 |
| 0.543069 |
| -0.36576 |
| -1.345465 |
| -0.577036 |
Note
与 pandas 的 Series 不同,Polars 的 Series 没有行索引。 Polars 是以列为中心的——数据访问是通过列表达式和布尔掩码来管理的,而不是通过行标签。 更多详情请参阅 面向 pandas 用户的 Polars 迁移指南。
Polars 的 Series 建立在 Apache Arrow 数组之上,并支持许多我们熟悉的操作
s * 100
| daily returns |
|---|
| f64 |
| 54.306927 |
| -36.575974 |
| -134.546478 |
| -57.703566 |
绝对值可以通过一个方法来获得
s.abs()
| daily returns |
|---|
| f64 |
| 0.543069 |
| 0.36576 |
| 1.345465 |
| 0.577036 |
我们也可以快速获取汇总统计信息
s.describe()
| statistic | value |
|---|---|
| str | f64 |
| "count" | 4.0 |
| "null_count" | 0.0 |
| "mean" | -0.436298 |
| "std" | 0.776858 |
| "min" | -1.345465 |
| "25%" | -0.577036 |
| "50%" | -0.36576 |
| "75%" | -0.36576 |
| "max" | 0.543069 |
由于 Polars 没有行索引,带标签的数据需要用 DataFrame 来表示。
例如,要将股票代码与收益关联起来:
df = pl.DataFrame({
'company': ['AMZN', 'AAPL', 'MSFT', 'GOOG'],
'daily returns': np.random.randn(4)
})
df
| company | daily returns |
|---|---|
| str | f64 |
| "AMZN" | 0.552575 |
| "AAPL" | 1.886849 |
| "MSFT" | -0.433338 |
| "GOOG" | 0.813842 |
我们通过对某一列进行过滤来访问某个值
df.filter(
pl.col('company') == 'AMZN'
).select('daily returns').item()
0.5525750956035862
更新数据同样使用表达式,而不是按索引赋值
df = df.with_columns(
pl.when(pl.col('company') == 'AMZN')
.then(0)
.otherwise(pl.col('daily returns'))
.alias('daily returns')
)
df
| company | daily returns |
|---|---|
| str | f64 |
| "AMZN" | 0.0 |
| "AAPL" | 1.886849 |
| "MSFT" | -0.433338 |
| "GOOG" | 0.813842 |
我们还可以检查成员关系
'AAPL' in df['company']
True
19.3. DataFrames#
Series 是单独一列数据,而 DataFrame 则包含多个列,每个变量对应一列。
和 Pandas 一样,我们用 Penn World Tables 中的数据来进行操作。
我们使用 pl.read_csv 来读取数据
url = ('https://raw.githubusercontent.com/QuantEcon/'
'lecture-python-programming/main/lectures/_static/'
'lecture_specific/pandas/data/test_pwt.csv')
df = pl.read_csv(url)
df
| country | country isocode | year | POP | XRAT | tcgdp | cc | cg |
|---|---|---|---|---|---|---|---|
| str | str | i64 | f64 | f64 | f64 | f64 | f64 |
| "Argentina" | "ARG" | 2000 | 37335.653 | 0.9995 | 295072.21869 | 75.716805 | 5.578804 |
| "Australia" | "AUS" | 2000 | 19053.186 | 1.72483 | 541804.6521 | 67.759026 | 6.720098 |
| "India" | "IND" | 2000 | 1.0063e6 | 44.9416 | 1.7281e6 | 64.575551 | 14.072206 |
| "Israel" | "ISR" | 2000 | 6114.57 | 4.07733 | 129253.89423 | 64.436451 | 10.266688 |
| "Malawi" | "MWI" | 2000 | 11801.505 | 59.543808 | 5026.221784 | 74.707624 | 11.658954 |
| "South Africa" | "ZAF" | 2000 | 45064.098 | 6.93983 | 227242.36949 | 72.71871 | 5.726546 |
| "United States" | "USA" | 2000 | 282171.957 | 1.0 | 9.8987e6 | 72.347054 | 6.032454 |
| "Uruguay" | "URY" | 2000 | 3219.793 | 12.099592 | 25255.961693 | 78.97874 | 5.108068 |
19.3.1. 选择数据#
我们可以通过切片来选择行,通过列名来选择列
df[2:5]
| country | country isocode | year | POP | XRAT | tcgdp | cc | cg |
|---|---|---|---|---|---|---|---|
| str | str | i64 | f64 | f64 | f64 | f64 | f64 |
| "India" | "IND" | 2000 | 1.0063e6 | 44.9416 | 1.7281e6 | 64.575551 | 14.072206 |
| "Israel" | "ISR" | 2000 | 6114.57 | 4.07733 | 129253.89423 | 64.436451 | 10.266688 |
| "Malawi" | "MWI" | 2000 | 11801.505 | 59.543808 | 5026.221784 | 74.707624 | 11.658954 |
要选择特定的列,可以向 select 传入一个名称列表
df.select(['country', 'tcgdp'])
| country | tcgdp |
|---|---|
| str | f64 |
| "Argentina" | 295072.21869 |
| "Australia" | 541804.6521 |
| "India" | 1.7281e6 |
| "Israel" | 129253.89423 |
| "Malawi" | 5026.221784 |
| "South Africa" | 227242.36949 |
| "United States" | 9.8987e6 |
| "Uruguay" | 25255.961693 |
这些操作也可以组合使用
df[2:5].select(['country', 'tcgdp'])
| country | tcgdp |
|---|---|
| str | f64 |
| "India" | 1.7281e6 |
| "Israel" | 129253.89423 |
| "Malawi" | 5026.221784 |
19.3.2. 按条件过滤#
filter 方法接受由 pl.col 构建的布尔表达式
df.filter(pl.col('POP') >= 20000)
| country | country isocode | year | POP | XRAT | tcgdp | cc | cg |
|---|---|---|---|---|---|---|---|
| str | str | i64 | f64 | f64 | f64 | f64 | f64 |
| "Argentina" | "ARG" | 2000 | 37335.653 | 0.9995 | 295072.21869 | 75.716805 | 5.578804 |
| "India" | "IND" | 2000 | 1.0063e6 | 44.9416 | 1.7281e6 | 64.575551 | 14.072206 |
| "South Africa" | "ZAF" | 2000 | 45064.098 | 6.93983 | 227242.36949 | 72.71871 | 5.726546 |
| "United States" | "USA" | 2000 | 282171.957 | 1.0 | 9.8987e6 | 72.347054 | 6.032454 |
可以使用 &(与)和 |(或)来组合多个条件
df.filter(
(pl.col('country').is_in(['Argentina', 'India', 'South Africa'])) &
(pl.col('POP') > 40000)
)
| country | country isocode | year | POP | XRAT | tcgdp | cc | cg |
|---|---|---|---|---|---|---|---|
| str | str | i64 | f64 | f64 | f64 | f64 | f64 |
| "India" | "IND" | 2000 | 1.0063e6 | 44.9416 | 1.7281e6 | 64.575551 | 14.072206 |
| "South Africa" | "ZAF" | 2000 | 45064.098 | 6.93983 | 227242.36949 | 72.71871 | 5.726546 |
表达式可以涉及跨列的算术运算
df.filter(
(pl.col('cc') + pl.col('cg') >= 80) & (pl.col('POP') <= 20000)
)
| country | country isocode | year | POP | XRAT | tcgdp | cc | cg |
|---|---|---|---|---|---|---|---|
| str | str | i64 | f64 | f64 | f64 | f64 | f64 |
| "Malawi" | "MWI" | 2000 | 11801.505 | 59.543808 | 5026.221784 | 74.707624 | 11.658954 |
| "Uruguay" | "URY" | 2000 | 3219.793 | 12.099592 | 25255.961693 | 78.97874 | 5.108068 |
选择家庭消费占比最大的国家
df.filter(pl.col('cc') == pl.col('cc').max())
| country | country isocode | year | POP | XRAT | tcgdp | cc | cg |
|---|---|---|---|---|---|---|---|
| str | str | i64 | f64 | f64 | f64 | f64 | f64 |
| "Uruguay" | "URY" | 2000 | 3219.793 | 12.099592 | 25255.961693 | 78.97874 | 5.108068 |
19.3.3. 列表达式#
与 pandas 的一个关键区别在于,Polars 使用列表达式来进行转换,而不是逐元素调用 apply。
下面是一个计算每个数值列最大值的示例
df.select(
pl.col(['year', 'POP', 'XRAT', 'tcgdp', 'cc', 'cg'])
.max()
.name.suffix('_max')
)
| year_max | POP_max | XRAT_max | tcgdp_max | cc_max | cg_max |
|---|---|---|---|---|---|
| i64 | f64 | f64 | f64 | f64 | f64 |
| 2000 | 1.0063e6 | 59.543808 | 9.8987e6 | 78.97874 | 14.072206 |
表达式可以在 with_columns 内部使用,用于添加或修改列
df.with_columns(
(pl.col('XRAT') / 10).alias('XRAT_scaled'),
pl.col(pl.Float64).round(2)
)
| country | country isocode | year | POP | XRAT | tcgdp | cc | cg | XRAT_scaled |
|---|---|---|---|---|---|---|---|---|
| str | str | i64 | f64 | f64 | f64 | f64 | f64 | f64 |
| "Argentina" | "ARG" | 2000 | 37335.65 | 1.0 | 295072.22 | 75.72 | 5.58 | 0.09995 |
| "Australia" | "AUS" | 2000 | 19053.19 | 1.72 | 541804.65 | 67.76 | 6.72 | 0.172483 |
| "India" | "IND" | 2000 | 1006300.3 | 44.94 | 1.7281e6 | 64.58 | 14.07 | 4.49416 |
| "Israel" | "ISR" | 2000 | 6114.57 | 4.08 | 129253.89 | 64.44 | 10.27 | 0.407733 |
| "Malawi" | "MWI" | 2000 | 11801.5 | 59.54 | 5026.22 | 74.71 | 11.66 | 5.954381 |
| "South Africa" | "ZAF" | 2000 | 45064.1 | 6.94 | 227242.37 | 72.72 | 5.73 | 0.693983 |
| "United States" | "USA" | 2000 | 282171.96 | 1.0 | 9.8987e6 | 72.35 | 6.03 | 0.1 |
| "Uruguay" | "URY" | 2000 | 3219.79 | 12.1 | 25255.96 | 78.98 | 5.11 | 1.209959 |
条件逻辑使用 pl.when(...).then(...).otherwise(...)
df.with_columns(
pl.when(pl.col('POP') >= 20000)
.then(pl.col('POP'))
.otherwise(None)
.alias('POP_filtered')
).select(['country', 'POP', 'POP_filtered'])
| country | POP | POP_filtered |
|---|---|---|
| str | f64 | f64 |
| "Argentina" | 37335.653 | 37335.653 |
| "Australia" | 19053.186 | null |
| "India" | 1.0063e6 | 1.0063e6 |
| "Israel" | 6114.57 | null |
| "Malawi" | 11801.505 | null |
| "South Africa" | 45064.098 | 45064.098 |
| "United States" | 282171.957 | 282171.957 |
| "Uruguay" | 3219.793 | null |
Note
Polars 提供了 map_elements 作为逐行应用任意 Python 函数的应急方案,
但它会绕过经过优化的表达式引擎,因此只要存在原生表达式,就应该避免使用它。
19.3.4. 缺失值#
让我们插入一些空值来演示插补技术
df_nulls = df.with_row_index().with_columns(
pl.when(pl.col('index') == 0)
.then(None).otherwise(pl.col('XRAT')).alias('XRAT'),
pl.when(pl.col('index') == 3)
.then(None).otherwise(pl.col('cc')).alias('cc'),
pl.when(pl.col('index') == 5)
.then(None).otherwise(pl.col('tcgdp')).alias('tcgdp'),
pl.when(pl.col('index') == 6)
.then(None).otherwise(pl.col('POP')).alias('POP'),
).drop('index')
df_nulls
| country | country isocode | year | POP | XRAT | tcgdp | cc | cg |
|---|---|---|---|---|---|---|---|
| str | str | i64 | f64 | f64 | f64 | f64 | f64 |
| "Argentina" | "ARG" | 2000 | 37335.653 | null | 295072.21869 | 75.716805 | 5.578804 |
| "Australia" | "AUS" | 2000 | 19053.186 | 1.72483 | 541804.6521 | 67.759026 | 6.720098 |
| "India" | "IND" | 2000 | 1.0063e6 | 44.9416 | 1.7281e6 | 64.575551 | 14.072206 |
| "Israel" | "ISR" | 2000 | 6114.57 | 4.07733 | 129253.89423 | null | 10.266688 |
| "Malawi" | "MWI" | 2000 | 11801.505 | 59.543808 | 5026.221784 | 74.707624 | 11.658954 |
| "South Africa" | "ZAF" | 2000 | 45064.098 | 6.93983 | null | 72.71871 | 5.726546 |
| "United States" | "USA" | 2000 | null | 1.0 | 9.8987e6 | 72.347054 | 6.032454 |
| "Uruguay" | "URY" | 2000 | 3219.793 | 12.099592 | 25255.961693 | 78.97874 | 5.108068 |
将所有空值填充为零
df_nulls.fill_null(0)
| country | country isocode | year | POP | XRAT | tcgdp | cc | cg |
|---|---|---|---|---|---|---|---|
| str | str | i64 | f64 | f64 | f64 | f64 | f64 |
| "Argentina" | "ARG" | 2000 | 37335.653 | 0.0 | 295072.21869 | 75.716805 | 5.578804 |
| "Australia" | "AUS" | 2000 | 19053.186 | 1.72483 | 541804.6521 | 67.759026 | 6.720098 |
| "India" | "IND" | 2000 | 1.0063e6 | 44.9416 | 1.7281e6 | 64.575551 | 14.072206 |
| "Israel" | "ISR" | 2000 | 6114.57 | 4.07733 | 129253.89423 | 0.0 | 10.266688 |
| "Malawi" | "MWI" | 2000 | 11801.505 | 59.543808 | 5026.221784 | 74.707624 | 11.658954 |
| "South Africa" | "ZAF" | 2000 | 45064.098 | 6.93983 | 0.0 | 72.71871 | 5.726546 |
| "United States" | "USA" | 2000 | 0.0 | 1.0 | 9.8987e6 | 72.347054 | 6.032454 |
| "Uruguay" | "URY" | 2000 | 3219.793 | 12.099592 | 25255.961693 | 78.97874 | 5.108068 |
或者用列的均值来填充
cols = ['cc', 'tcgdp', 'POP', 'XRAT']
df_nulls.with_columns(
pl.col(cols).fill_null(pl.col(cols).mean())
)
| country | country isocode | year | POP | XRAT | tcgdp | cc | cg |
|---|---|---|---|---|---|---|---|
| str | str | i64 | f64 | f64 | f64 | f64 | f64 |
| "Argentina" | "ARG" | 2000 | 37335.653 | 18.618141 | 295072.21869 | 75.716805 | 5.578804 |
| "Australia" | "AUS" | 2000 | 19053.186 | 1.72483 | 541804.6521 | 67.759026 | 6.720098 |
| "India" | "IND" | 2000 | 1.0063e6 | 44.9416 | 1.7281e6 | 64.575551 | 14.072206 |
| "Israel" | "ISR" | 2000 | 6114.57 | 4.07733 | 129253.89423 | 72.400502 | 10.266688 |
| "Malawi" | "MWI" | 2000 | 11801.505 | 59.543808 | 5026.221784 | 74.707624 | 11.658954 |
| "South Africa" | "ZAF" | 2000 | 45064.098 | 6.93983 | 1.8033e6 | 72.71871 | 5.726546 |
| "United States" | "USA" | 2000 | 161269.871714 | 1.0 | 9.8987e6 | 72.347054 | 6.032454 |
| "Uruguay" | "URY" | 2000 | 3219.793 | 12.099592 | 25255.961693 | 78.97874 | 5.108068 |
Polars 还支持向前填充(fill_null(strategy='forward'))以及插值。
在 scikit-learn 中还有更多高级插补工具可供使用。
19.3.5. 可视化#
让我们构建一个人均 GDP 列,并绘制出来
df = (df
.select(['country', 'POP', 'tcgdp'])
.rename({'POP': 'population', 'tcgdp': 'total GDP'})
.with_columns(
(pl.col('population') * 1e3).alias('population')
)
.with_columns(
(pl.col('total GDP') * 1e6 / pl.col('population'))
.alias('GDP percap')
)
.sort('GDP percap', descending=True)
)
df
| country | population | total GDP | GDP percap |
|---|---|---|---|
| str | f64 | f64 | f64 |
| "United States" | 2.82171957e8 | 9.8987e6 | 35080.381854 |
| "Australia" | 1.9053186e7 | 541804.6521 | 28436.433261 |
| "Israel" | 6.11457e6 | 129253.89423 | 21138.672749 |
| "Argentina" | 3.7335653e7 | 295072.21869 | 7903.229085 |
| "Uruguay" | 3.219793e6 | 25255.961693 | 7843.97062 |
| "South Africa" | 4.5064098e7 | 227242.36949 | 5042.647686 |
| "India" | 1.0063e9 | 1.7281e6 | 1717.324719 |
| "Malawi" | 1.1801505e7 | 5026.221784 | 425.896679 |
我们可以直接提取列用于 matplotlib
Note
Polars 还提供了基于 Altair 的内置绘图 API
(例如 df.plot.bar(x=..., y=...))。
为了与本系列讲座的其他部分保持一致,我们在这里使用 matplotlib。
fig, ax = plt.subplots()
ax.bar(df['country'].to_list(), df['GDP percap'].to_list())
ax.set_xlabel('country', fontsize=12)
ax.set_ylabel('GDP per capita', fontsize=12)
plt.xticks(rotation=45, ha='right')
plt.tight_layout()
plt.show()
19.4. 惰性求值#
Polars 最强大的特性之一是惰性求值(lazy evaluation)。
惰性模式并不会立即执行每一个操作,而是先收集完整的查询计划,再在运行前对其进行优化。
19.4.1. 立即执行与惰性执行#
# 重新加载数据集
url = ('https://raw.githubusercontent.com/QuantEcon/'
'lecture-python-programming/main/lectures/_static/'
'lecture_specific/pandas/data/test_pwt.csv')
df_full = pl.read_csv(url)
立即执行(eager) API 会立即执行(就像 pandas 那样)
result_eager = (df_full
.filter(pl.col('tcgdp') > 1000)
.select(['country', 'year', 'tcgdp'])
.sort('tcgdp', descending=True)
)
result_eager.head()
| country | year | tcgdp |
|---|---|---|
| str | i64 | f64 |
| "United States" | 2000 | 9.8987e6 |
| "India" | 2000 | 1.7281e6 |
| "Australia" | 2000 | 541804.6521 |
| "Argentina" | 2000 | 295072.21869 |
| "South Africa" | 2000 | 227242.36949 |
惰性(lazy) API 则会构建一个查询计划
lazy_query = (df_full.lazy()
.filter(pl.col('tcgdp') > 1000)
.select(['country', 'year', 'tcgdp'])
.sort('tcgdp', descending=True)
)
print(lazy_query.explain())
SORT BY [descending: [true]] [col("tcgdp")]
FILTER col("tcgdp") > 1000.0
FROM
DF ["country", "country isocode", "year", "POP", ...]; PROJECT["country", "year", "tcgdp"] 3/8 COLUMNS
调用 collect 来执行该计划
result_lazy = lazy_query.collect()
result_lazy.head()
| country | year | tcgdp |
|---|---|---|
| str | i64 | f64 |
| "United States" | 2000 | 9.8987e6 |
| "India" | 2000 | 1.7281e6 |
| "Australia" | 2000 | 541804.6521 |
| "Argentina" | 2000 | 295072.21869 |
| "South Africa" | 2000 | 227242.36949 |
19.4.2. 查询优化#
惰性执行引擎会自动应用若干优化:
谓词下推——过滤操作会被尽可能提前应用
投影下推——只从数据源读取所需的列
公共子表达式消除——重复的计算会被合并
让我们看看 Polars 是如何重写一个多步骤查询的
optimized = (df_full.lazy()
.select(['country', 'year', 'tcgdp', 'POP'])
.filter(pl.col('tcgdp') > 500)
.with_columns(
(pl.col('tcgdp') / pl.col('POP')).alias('gdp_per_capita')
)
.filter(pl.col('gdp_per_capita') > 10)
.select(['country', 'year', 'gdp_per_capita'])
)
print("Optimized plan:")
print(optimized.explain())
Optimized plan:
FILTER col("gdp_per_capita") > 10.0
FROM
simple π 3/3 ["country", "year", ... 1 other column]
WITH_COLUMNS:
[(col("tcgdp") / col("POP")).alias("gdp_per_capita")]
FILTER col("tcgdp") > 500.0
FROM
DF ["country", "country isocode", "year", "POP", ...]; PROJECT["country", "year", "tcgdp", "POP"] 4/8 COLUMNS
执行该计划便可得到最终结果
optimized.collect()
| country | year | gdp_per_capita |
|---|---|---|
| str | i64 | f64 |
| "Australia" | 2000 | 28.436433 |
| "Israel" | 2000 | 21.138673 |
| "United States" | 2000 | 35.080382 |
19.4.3. 性能比较#
让我们在同一任务上比较 pandas、Polars 立即执行模式和 Polars 惰性执行模式。
我们先从一个小数据集(上面使用过的 Penn World Tables 数据)开始,说明对于小数据而言,各方案之间的差异微不足道
import pandas as pd
import time
# 小数据集 -- Penn World Tables(约 8 行)
url = ('https://raw.githubusercontent.com/QuantEcon/'
'lecture-python-programming/main/lectures/_static/'
'lecture_specific/pandas/data/test_pwt.csv')
small_pd = pd.read_csv(url)
small_pl = pl.read_csv(url)
现在我们对每个库中相同的过滤-选择-排序操作进行计时
# pandas
start = time.perf_counter()
_ = (small_pd
.query('tcgdp > 500')
[['country', 'year', 'tcgdp', 'POP']]
.assign(gdp_pc=lambda d: d['tcgdp'] / d['POP'])
.sort_values('gdp_pc', ascending=False))
pd_small = time.perf_counter() - start
# Polars 立即执行模式
start = time.perf_counter()
_ = (small_pl
.filter(pl.col('tcgdp') > 500)
.select(['country', 'year', 'tcgdp', 'POP'])
.with_columns((pl.col('tcgdp') / pl.col('POP')).alias('gdp_pc'))
.sort('gdp_pc', descending=True))
pl_small = time.perf_counter() - start
print(f"Small data -- pandas: {pd_small:.4f}s | Polars eager: {pl_small:.4f}s")
Small data -- pandas: 0.0043s | Polars eager: 0.0010s
对于寥寥数行的数据,速度差异并不重要——你可以选择自己更习惯使用的 API。
现在让我们把数据规模扩大到 500 万行,此时差异就会变得明显。
我们的任务是:过滤出 value > 0 的行,计算加权乘积
value * weight,然后对每个分组内该乘积取均值——
也就是一个分组加权平均。
n = 5_000_000
np.random.seed(42)
groups = np.random.choice(['A', 'B', 'C', 'D'], n)
values = np.random.randn(n)
weights = np.random.rand(n)
extra1 = np.random.randn(n)
extra2 = np.random.randn(n)
big_pd = pd.DataFrame({
'group': groups, 'value': values,
'weight': weights, 'extra1': extra1, 'extra2': extra2
})
big_pl = pl.DataFrame({
'group': groups, 'value': values,
'weight': weights, 'extra1': extra1, 'extra2': extra2
})
首先是 pandas 基准测试
start = time.perf_counter()
tmp = big_pd[big_pd['value'] > 0][['group', 'value', 'weight']].copy()
tmp['weighted'] = tmp['value'] * tmp['weight']
_ = tmp.groupby('group')['weighted'].mean()
pd_time = time.perf_counter() - start
print(f"pandas: {pd_time:.4f}s")
pandas: 0.1706s
接下来是 Polars 的立即执行模式
start = time.perf_counter()
_ = (big_pl
.filter(pl.col('value') > 0)
.select(['group', 'value', 'weight'])
.with_columns(
(pl.col('value') * pl.col('weight')).alias('weighted'))
.group_by('group')
.agg(pl.col('weighted').mean()))
eager_time = time.perf_counter() - start
print(f"Polars eager: {eager_time:.4f}s")
Polars eager: 0.0354s
最后是 Polars 的惰性执行模式
start = time.perf_counter()
_ = (big_pl.lazy()
.filter(pl.col('value') > 0)
.select(['group', 'value', 'weight'])
.with_columns(
(pl.col('value') * pl.col('weight')).alias('weighted'))
.group_by('group')
.agg(pl.col('weighted').mean())
.collect())
lazy_time = time.perf_counter() - start
print(f"Polars lazy: {lazy_time:.4f}s")
Polars lazy: 0.0309s
由此可以得出以下结论:
对于小数据(数千行),pandas 与 Polars 的表现相似—— 可以根据 API 偏好和生态系统兼容性来选择。
对于中大型数据(数十万行及以上),得益于其 Rust 引擎、 并行执行以及(在惰性模式下的)查询优化,Polars 可以显著更快。
当从磁盘读取数据时,惰性 API 尤其强大——scan_csv 会直接返回一个 LazyFrame,因此过滤和投影操作会被下推到文件读取器中执行。
Tip
在处理大型 CSV 文件时,使用 pl.scan_csv(path) 而不是 pl.read_csv(path)。
这样只有你实际需要的列和行才会从磁盘中被读取。
参见 Polars I/O 文档。
19.5. 在线数据源#
与 Pandas 一样,Python 可以非常方便地查询在线数据库。
对经济学家而言,一个重要的数据库是 FRED——由圣路易斯联储维护的庞大时间序列数据集。
Polars 的 read_csv 可以直接从 URL 获取数据。
我们使用 try_parse_dates=True 来自动解析日期列
fred_url = ('https://fred.stlouisfed.org/graph/fredgraph.csv?'
'bgcolor=%23e1e9f0&chart_type=line&drp=0&'
'fo=open%20sans&graph_bgcolor=%23ffffff&'
'height=450&mode=fred&recession_bars=on&'
'txtcolor=%23444444&ts=12&tts=12&width=1318&'
'nt=0&thu=0&trc=0&show_legend=yes&'
'show_axis_titles=yes&show_tooltip=yes&'
'id=UNRATE&scale=left&cosd=1948-01-01&'
'coed=2024-06-01&line_color=%234572a7&'
'link_values=false&line_style=solid&'
'mark_type=none&mw=3&lw=2&ost=-99999&'
'oet=99999&mma=0&fml=a&fq=Monthly&fam=avg&'
'fgst=lin&fgsnd=2020-02-01&line_index=1&'
'transformation=lin&vintage_date=2024-07-29&'
'revision_date=2024-07-29&nd=1948-01-01')
data = pl.read_csv(fred_url, try_parse_dates=True)
让我们查看前几行
data.head()
| observation_date | UNRATE |
|---|---|
| date | f64 |
| 1948-01-01 | 3.4 |
| 1948-02-01 | 3.8 |
| 1948-03-01 | 4.0 |
| 1948-04-01 | 3.9 |
| 1948-05-01 | 3.5 |
以及获取汇总统计信息
data.describe()
| statistic | observation_date | UNRATE |
|---|---|---|
| str | str | f64 |
| "count" | "918" | 918.0 |
| "null_count" | "0" | 0.0 |
| "mean" | "1986-03-17 06:30:35.294117" | 5.693246 |
| "std" | null | 1.710248 |
| "min" | "1948-01-01" | 2.5 |
| "25%" | "1967-02-01" | 4.4 |
| "50%" | "1986-04-01" | 5.5 |
| "75%" | "2005-05-01" | 6.7 |
| "max" | "2024-06-01" | 14.8 |
绘制 2006 年到 2012 年的失业率
filtered = data.filter(
(pl.col('observation_date') >= pl.date(2006, 1, 1)) &
(pl.col('observation_date') <= pl.date(2012, 12, 31))
)
fig, ax = plt.subplots()
ax.plot(filtered['observation_date'].to_list(),
filtered['UNRATE'].to_list())
ax.set_title('US Unemployment Rate')
ax.set_xlabel('year', fontsize=12)
ax.set_ylabel('%', fontsize=12)
plt.show()
Polars 支持多种文件格式,包括 Excel、JSON、Parquet,以及直接连接数据库。
19.6. 练习#
Exercise 19.1
使用以下导入:
import datetime as dt
import yfinance as yf
编写一个程序,计算以下几只股票在 2021 年间价格的百分比变化:
ticker_list = {'INTC': 'Intel',
'MSFT': 'Microsoft',
'IBM': 'IBM',
'BHP': 'BHP',
'TM': 'Toyota',
'AAPL': 'Apple',
'AMZN': 'Amazon',
'C': 'Citigroup',
'QCOM': 'Qualcomm',
'KO': 'Coca-Cola',
'GOOG': 'Google'}
下面是一个将收盘价读入 Polars DataFrame 的函数:
def read_data_polars(ticker_list,
start=dt.datetime(2021, 1, 1),
end=dt.datetime(2021, 12, 31)):
"""
Read closing price data from Yahoo Finance
and return a Polars DataFrame.
"""
dataframes = []
for tick in ticker_list:
stock = yf.Ticker(tick)
prices = stock.history(start=start, end=end)
df = pl.DataFrame({
'Date': list(prices.index.date),
tick: prices['Close'].values
}).with_columns(pl.col('Date').cast(pl.Date))
dataframes.append(df)
result = dataframes[0]
for df in dataframes[1:]:
result = result.join(
df, on='Date', how='full', coalesce=True
)
return result.sort('Date')
ticker = read_data_polars(ticker_list)
Note
Polars 的连接(join)操作不保证输出行的顺序——
只匹配一侧的键值会被追加,而不是被插入到相应位置。
这和前面提到的”没有索引、没有自动对齐”的主题是一致的:
由于没有行标签可供对齐,排序需要我们显式地去请求。
因此在返回结果之前会调用 sort('Date'),
后续任何 first()/last() 计算都依赖于这一排序结果。
请补充完整该程序,将结果绘制为柱状图。
Solution
使用 Polars 表达式计算百分比变化:
price_change = ticker.select([
((pl.col(tick).last() / pl.col(tick).first() - 1) * 100)
.alias(tick)
for tick in ticker_list.keys()
]).transpose(
include_header=True,
header_name='ticker',
column_names=['pct_change']
).with_columns(
pl.col('ticker')
.replace_strict(ticker_list, default=pl.col('ticker'))
.alias('company')
).sort('pct_change')
print(price_change)
shape: (11, 3)
┌────────┬────────────┬───────────┐
│ ticker ┆ pct_change ┆ company │
│ --- ┆ --- ┆ --- │
│ str ┆ f64 ┆ str │
╞════════╪════════════╪═══════════╡
│ BHP ┆ -2.249085 ┆ BHP │
│ C ┆ 3.550569 ┆ Citigroup │
│ AMZN ┆ 5.845049 ┆ Amazon │
│ INTC ┆ 6.868565 ┆ Intel │
│ KO ┆ 14.922487 ┆ Coca-Cola │
│ … ┆ … ┆ … │
│ TM ┆ 23.41675 ┆ Toyota │
│ QCOM ┆ 25.318529 ┆ Qualcomm │
│ AAPL ┆ 38.550754 ┆ Apple │
│ MSFT ┆ 57.179604 ┆ Microsoft │
│ GOOG ┆ 68.960882 ┆ Google │
└────────┴────────────┴───────────┘
直接使用 matplotlib 绘制结果:
companies = price_change['company'].to_list()
changes = price_change['pct_change'].to_list()
colors = ['red' if x < 0 else 'blue' for x in changes]
fig, ax = plt.subplots(figsize=(10, 8))
ax.bar(companies, changes, color=colors)
ax.set_xlabel('stock', fontsize=12)
ax.set_ylabel('percentage change in price', fontsize=12)
plt.xticks(rotation=45, ha='right')
plt.tight_layout()
plt.show()
Exercise 19.2
使用 Exercise 19.1 中的 read_data_polars,求出以下指数的同比百分比变化:
indices_list = {'^GSPC': 'S&P 500',
'^IXIC': 'NASDAQ',
'^DJI': 'Dow Jones',
'^N225': 'Nikkei'}
将结果绘制为时间序列图。
Solution
indices_data = read_data_polars(
indices_list,
start=dt.datetime(1971, 1, 1),
end=dt.datetime(2021, 12, 31)
)
indices_data = indices_data.with_columns(
pl.col('Date').dt.year().alias('year')
)
使用分组操作计算逐年收益:
yearly_returns = indices_data.group_by('year').agg([
*[pl.col(idx).drop_nulls().first().alias(f'{idx}_first')
for idx in indices_list],
*[pl.col(idx).drop_nulls().last().alias(f'{idx}_last')
for idx in indices_list]
])
for idx, name in indices_list.items():
yearly_returns = yearly_returns.with_columns(
((pl.col(f'{idx}_last') - pl.col(f'{idx}_first'))
/ pl.col(f'{idx}_first') * 100).alias(name)
)
yearly_returns = (yearly_returns
.select(['year', *indices_list.values()])
.sort('year')
)
print(yearly_returns)
shape: (51, 5)
┌──────┬────────────┬────────────┬───────────┬────────────┐
│ year ┆ S&P 500 ┆ NASDAQ ┆ Dow Jones ┆ Nikkei │
│ --- ┆ --- ┆ --- ┆ --- ┆ --- │
│ i32 ┆ f64 ┆ f64 ┆ f64 ┆ f64 │
╞══════╪════════════╪════════════╪═══════════╪════════════╡
│ 1971 ┆ 12.002188 ┆ 14.120003 ┆ null ┆ 36.407234 │
│ 1972 ┆ 16.110952 ┆ 17.668274 ┆ null ┆ 92.011231 │
│ 1973 ┆ -18.094035 ┆ -31.523435 ┆ null ┆ -17.697016 │
│ 1974 ┆ -29.811633 ┆ -35.350697 ┆ null ┆ -9.914309 │
│ 1975 ┆ 28.4209 ┆ 27.874797 ┆ null ┆ 16.798024 │
│ … ┆ … ┆ … ┆ … ┆ … │
│ 2017 ┆ 18.415027 ┆ 27.155799 ┆ 24.331151 ┆ 16.182267 │
│ 2018 ┆ -7.009394 ┆ -5.303631 ┆ -6.028635 ┆ -14.853703 │
│ 2019 ┆ 28.714796 ┆ 34.603667 ┆ 22.23998 ┆ 20.931737 │
│ 2020 ┆ 15.292907 ┆ 41.751104 ┆ 6.019231 ┆ 18.269064 │
│ 2021 ┆ 29.132182 ┆ 23.964416 ┆ 20.428169 ┆ 5.625169 │
└──────┴────────────┴────────────┴───────────┴────────────┘
汇总统计信息:
yearly_returns.select(list(indices_list.values())).describe()
| statistic | S&P 500 | NASDAQ | Dow Jones | Nikkei |
|---|---|---|---|---|
| str | f64 | f64 | f64 | f64 |
| "count" | 51.0 | 51.0 | 30.0 | 51.0 |
| "null_count" | 0.0 | 0.0 | 21.0 | 0.0 |
| "mean" | 9.20986 | 13.094786 | 9.10453 | 7.850346 |
| "std" | 16.398012 | 24.616625 | 14.134825 | 24.384181 |
| "min" | -37.58465 | -40.197764 | -32.716831 | -39.695649 |
| "25%" | 0.255632 | 2.470561 | 2.072082 | -6.094919 |
| "50%" | 11.677594 | 14.312575 | 9.387316 | 7.667034 |
| "75%" | 19.671602 | 27.874797 | 21.451016 | 20.931737 |
| "max" | 34.157394 | 84.294285 | 33.311106 | 92.011231 |
将每个指数绘制在一个子图中:
fig, axes = plt.subplots(2, 2, figsize=(12, 10))
years = yearly_returns['year'].to_list()
for iter_, ax in enumerate(axes.flatten()):
name = list(indices_list.values())[iter_]
values = yearly_returns[name].to_list()
ax.plot(years, values, 'o-', linewidth=2, markersize=4)
ax.axhline(y=0, color='k', linestyle='--', alpha=0.3)
ax.set_ylabel('yearly return (%)', fontsize=12)
ax.set_xlabel('year', fontsize=12)
ax.set_title(name, fontsize=12)
plt.tight_layout()
plt.show()