春江暮客

春江暮客的个人学习分享网站

Python 大型 CSV 随机抽样:用固定种子实现蓄水池算法

2026-10-08 技术
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 放入内存的情况下选出少量记录。按完整记录解析,保存随机种子和输入标识,并在使用输出前验证整个数据流。如果需要标签平衡、唯一实体抽样,或隔离相关分组,应另行设计抽样规则。

标签

1024 12306 ablang adsense agents.md ai ai-agent ai-agents ai-seo algorithm amp antibodies automation bioinformatics blockchain boltz bootstrapping boxes c-index cca cdn chatgpt cli cloudflare codex cofoldarena copy cpu监控 csv cuda curl data-leakage data-processing data-quality data-validation datascience datavisualization deployment desktop-app devtools disown docker dovecot download electron esm2 esm3 esmc esmfold2 faceswap fasta fastmcp ffmpeg file-io flashppi flask folium frontend game generator git github-actions google grep hls html http hugo indexnow javascript jev json just k-means kaggle langfuse leecode linux list litellm llm llms.txt logging logs lollipop m3u8 machine-learning macos manacher matplotlib mcp mirror model-evaluation mp3 mp4 mpnn multiomics mutation mysql nanobert nanobodies networkx nginx normalize ollama omegatherm pandas password pep-723 phaser pillow pip postfix preprocessing print protein-design protein-interactions protein-language-models protein-stability protein-structure proxy pydantic pyecharts pyqt python python3 r raincloud reproducibility requests reservoir-sampling rfoptimization rg ripgrep rosettafold3 roundcube rsync s-tui sampling scale scikit-learn scrapy screen seaborn security selenium seo sha256 sklearn solana somaticsignatures spl sqlite ssh standardize static-site subprocess system-one tensorflow tkinter tron tronpy turtle typesafe-ai usdt uv vhhbert vite webp wordcloud wordpress workflow yaml 后台 寓言 概率 经济 贸易 迅雷解析 钱包

友情链接

其它