Python 大型 CSV 随机抽样:用固定种子实现蓄水池算法
你需要从一个无法整体放入内存的 CSV 文件中随机抽取少量记录。直接取前一千行,可能偏向排在前面的数据来源、标签或日期。蓄水池抽样可以在读取文件一次的过程中选出固定数量的记录。
本文用 Python 标准库实现一个命令行工具:保留表头,按完整 CSV 记录抽样,并按原始顺序输出选中的记录。它衔接此前的大型 CSV 流式读取与验证教程。示例和七项测试已在 Python 3.14.7 上运行;本文没有进行大文件性能基准测试。
1. 先确定抽样单位
抽样单位是数据记录,不包含表头。输入有 N 条记录、要求抽取 k 条时,输出包含 min(k, N) 条记录。N 不小于 k 时,在均匀随机抽取的前提下,每条记录被选中的概率为 k/N。脚本使用基础蓄水池算法 Algorithm R,可参见 Vitter 的随机抽样论文。
| 输入或设置 | 行为 |
|---|---|
| UTF-8 编码、逗号分隔的 CSV,可带 BOM | 接受 |
| 列名非空且不重复的表头 | 必须提供 |
| 引号内的逗号和换行 | 作为完整记录的一部分读取 |
| 数据记录字段数不正确 | 报错,即使该记录最终不会被抽中 |
| 重复 ID 或重复记录 | 分别获得抽样机会 |
| 只有表头的文件 | 输出表头,报告抽中零条记录 |
| 输出路径已存在 | 拒绝覆盖 |
脚本将字段保留为字符串,因此 001 不会变成 1。它检查表格形状,不判断标签或序列的科学含义。真正的空白记录会因字段数不匹配而报错;正常记录中的空单元格则允许存在。
2. 下载并试跑小样例
使用 Python 3.11 或更高版本,无需第三方包。在新目录下载 sample_csv.py:
mkdir csv-sampling-demo
cd csv-sampling-demo
curl -fL https://www.bobobk.com/downloads/reservoir/sample_csv.py \
-o sample_csv.py
python3 sample_csv.py --help
创建下面这个虚构的六条记录数据表,抽取三条:
cat > input.csv <<'CSV'
sample_id,label
001,1
002,0
003,1
004,0
005,unknown
006,1
CSV
python3 sample_csv.py input.csv sample.csv --size 3 --seed 42
命令会输出以下 JSON 摘要:
{
"input_records": 6,
"requested_records": 3,
"sampled_records": 3,
"seed": 42
}
在本文测试的 Python 版本中,sample.csv 内容为:
sample_id,label
002,0
003,1
005,unknown
输出文件会按标准 CSV 引号规则和记录结束符重新写入。字段值保持不变,但不会逐字节保留原文件的排版。
3. 完整抽样脚本
如果希望检查并复制实现,将下面的代码保存为 sample_csv.py:
#!/usr/bin/env python3
"""Sample CSV records in one pass, then write a new output file."""
import argparse
import csv
import json
from pathlib import Path
import random
def sample_csv(source, size, seed):
if size < 1:
raise ValueError("sample size must be positive")
rng = random.Random(seed)
reservoir = []
seen = 0
with source.open(encoding="utf-8-sig", newline="") as handle:
reader = csv.reader(handle, strict=True)
header = next(reader, None)
if not header or any(not name.strip() for name in header):
raise ValueError("a header with nonempty column names is required")
if len(set(header)) != len(header):
raise ValueError("duplicate column names")
for seen, row in enumerate(reader, start=1):
if len(row) != len(header):
raise ValueError(
f"record {seen}: expected {len(header)} fields, got {len(row)}"
)
item = (seen, row)
if seen <= size:
reservoir.append(item)
else:
slot = rng.randrange(seen)
if slot < size:
reservoir[slot] = item
reservoir.sort(key=lambda item: item[0])
return header, [row for _, row in reservoir], seen
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("input", type=Path)
parser.add_argument("output", type=Path)
parser.add_argument("--size", required=True, type=int)
parser.add_argument("--seed", type=int, default=42)
args = parser.parse_args()
if args.size < 1:
parser.error("--size must be positive")
try:
header, rows, total = sample_csv(args.input, args.size, args.seed)
# Exclusive creation refuses to overwrite any existing destination.
with args.output.open("x", encoding="utf-8", newline="") as handle:
writer = csv.writer(handle)
writer.writerow(header)
writer.writerows(rows)
except (OSError, UnicodeError, csv.Error, ValueError) as exc:
parser.exit(1, f"ERROR: {exc}\n")
print(json.dumps({
"input_records": total,
"requested_records": args.size,
"sampled_records": len(rows),
"seed": args.seed,
}, indent=2))
if __name__ == "__main__":
main()
csv.reader 负责识别记录边界,包括引号内的换行。文件使用 newline="" 打开,符合 Python CSV 文档的建议。按逗号拆字符串或按物理行抽样,都会破坏这类记录。
脚本先读完并检查输入,再打开输出文件。因此,解析错误或字段数错误不会创建目标文件。独占创建模式 "x" 防止覆盖已有文件,也包括输入文件本身。写入失败或进程中断仍可能留下不完整的新文件;后续步骤只能使用成功退出的命令生成的文件。
4. 为什么后面的记录仍有机会
前 k 条记录先填满蓄水池。此后读取第 i 条记录时,从 0 到 i−1 均匀抽取一个整数。如果小于 k,就替换对应位置;否则保留当前样本。
新记录进入样本的概率是 k/i。此前已选中的某条记录,在这一步保留下来的概率是 1−1/i;乘以它此前的入选概率 k/(i−1),结果同样为 k/i。最后的排序只恢复输入顺序,不改变入选记录。
程序扫描全部 N 条记录,保留至多 k 条样本,另加当前记录和解析缓冲区。因此,内存用量也取决于记录大小,而不只取决于条数。对于宽度有界的记录,扫描为 O(N),最后排序为 O(k log k),样本存储为 O(k)。减小 k 可以降低内存占用,但仍需读取整个文件。
5. 用大文件前,先验证可重复性
使用同一输入再次运行,并换一个输出文件名:
python3 sample_csv.py input.csv sample-repeat.csv --size 3 --seed 42
python3 - <<'PY'
import csv
from pathlib import Path
assert Path("sample.csv").read_bytes() == Path("sample-repeat.csv").read_bytes()
with open("sample.csv", encoding="utf-8", newline="") as handle:
rows = list(csv.reader(handle))
assert rows[0] == ["sample_id", "label"]
assert len(rows) - 1 == 3
print("PASS: repeatable sample, header and record count")
PY
最后应显示:
PASS: repeatable sample, header and record count
对于你已有的 large.csv,选择一个新输出路径:
python3 sample_csv.py large.csv sample-1000.csv --size 1000 --seed 42
python3 --version
固定随机种子有助于在输入内容、输入顺序、代码和 Python 环境不变时复现结果。脚本使用独立的 random.Random 实例。不要假定所有未来 Python 版本中的 randrange 都会产生完全相同的样本;随机数模块的可重复性说明区分了稳定保证和可能变化的算法。
记录种子、抽样数量、脚本版本和 Python 版本。用 SHA-256 清单标识输入,并保存生成的样本。同样的记录换一种排列顺序,抽样结果也可能改变。对于小数据集,换一个种子也不保证一定得到不同的样本。
6. 常见错误与适用范围
创建一条缺少字段的记录,观察失败行为:
cat > bad.csv <<'CSV'
sample_id,label
001,1
002
CSV
python3 sample_csv.py bad.csv bad-sample.csv --size 1 --seed 42
命令以非零状态退出,并报告:
ERROR: record 2: expected 2 fields, got 1
这里的记录编号不包含表头,按逻辑 CSV 记录计数,不是物理行号。即使设置 --size 1,第二条错误记录仍然会被检查。
| 问题 | 处理方式 |
|---|---|
| 输出文件已存在 | 换一个新的输出文件名 |
| 空文件或表头无效 | 提供列名非空且不重复的表头 |
| 字段数不一致 | 检查该记录的引号、分隔符和缺失单元格 |
| UTF-8 解码失败 | 确定源编码并转换后,再执行抽样 |
| 字段超过解析器限制 | 先检查源字段,确认合法后在脚本中设置合适的 csv.field_size_limit |
| 稀有标签抽中数量太少 | 使用明确设计的分层抽样流程 |
均匀抽取记录不会自动平衡标签。某个 ID 出现十次,它的十条记录就分别获得抽样机会;这不是对唯一实体抽样。如果目标单位是 ID,先检查 CSV 重复 ID。模型评估中,如果相关序列或重复测量必须放在同一组,应使用按组划分数据集的方法。
7. 总结
蓄水池抽样可以在不把整个 CSV 放入内存的情况下选出少量记录。按完整记录解析,保存随机种子和输入标识,并在使用输出前验证整个数据流。如果需要标签平衡、唯一实体抽样,或隔离相关分组,应另行设计抽样规则。
- 原文作者:春江暮客
- 原文链接:https://www.bobobk.com/python-csv-reservoir-sampling.html
- 版权声明:本作品采用 知识共享署名-非商业性使用-禁止演绎 4.0 国际许可协议 进行许可,非商业转载请注明出处(作者,原文链接),商业转载请联系作者获得授权。