本文是《算法工程师的工具箱:从一个想法到一次能跑的实验》系列的第 1 篇(共六篇)。下一篇:科学计算栈——NumPy 的形状直觉、Pandas 的错误分析、Matplotlib 的曲线

算法工程师每天写的 Python 不多,读的很多:transformersTrainer、别人的数据脚本、论文附带的训练代码。这些代码里反复出现同一小撮语法——__getitem__yield@torch.no_grad()with torch.autocast(...)**kwargs@dataclass——它们不是 PyTorch 发明的,是 Python 的协议:PyTorch 只是约定”你按这个形状写,我就能用”。本篇从要做的四件事出发——过一遍语料、写配置、读懂训练代码、把预处理跑快——把这一小撮语法带出来,讲到会读会用为止。Python 的机制(这些语法在解释器里怎么实现、GIL 是什么、装饰器怎么改函数)是 Infra 地图 01 系列的内容,那是本篇的深入篇,本篇末尾给出对照表。

全篇的核心问题是:

别人的训练代码里 __getitem__yield@torch.no_grad()with autocast(...)**kwargs 各在干什么?1 一份 10 GB 的 JSONL 语料怎么在 16 GB 内存的机器上过一遍、去重、统计?2 预处理太慢,开多线程为什么没用、开多进程为什么也到不了 8 倍?3

一、总览

1. 四件事、一小撮语法

要做的事 用到的语法
过一遍语料:读、过滤、去重、统计 with openjson生成器yield)、Counterpathlib
写一份能存能改的配置 @dataclassreplaceasdict、类型标注
读懂训练代码 __len__ / __getitem__、迭代协议、__call__装饰器上下文管理器*args / **kwargs
把预处理跑快 multiprocessing.Pool、GIL、chunksize
出错时定位 读 traceback、assert 形状、breakpoint()

2. 本文的章节安排

主题 内容
环境 一个项目一个环境;import torch 报错多半是环境的事
流式过一遍语料 一次读进内存 vs 生成器:19 MB 文件 103 MB vs 12 MB;生成器串成流水线
配置:dataclass 默认值、覆盖、存盘、构造时校验;可变默认值的坑
训练代码里的六个语法 PyTorch 的每个 API 对应 Python 的哪个协议;一个 40 行的”玩具 PyTorch”
多进程预处理 串行 / 线程 / 进程:0.46 s / 0.46 s / 0.15 s;为什么线程不行、进程也到不了 8×
出错的时候 traceback 从下往上读;三类最常见的错
越过哪条线进 Infra 01 本篇每一节的机制在 01 系列哪一篇
本文小结  
自测 五道题

配套脚本:00_python_in_use.py,只用标准库,文中的数字都来自它。

二、环境

一条规则:一个项目一个环境,环境里装什么写在文件里。

uv venv && source .venv/bin/activate              # 或 python -m venv .venv / conda create -n proj python=3.12
uv pip install torch --index-url https://download.pytorch.org/whl/cu124   # GPU 版 torch 要指定 CUDA 版本的源
uv pip install -r requirements.txt                # numpy pandas transformers ...
python -m pip list | grep -i torch                # python -m:用"当前这个 python"的 pip,不会装错环境

import torchModuleNotFoundErrortorch.cuda.is_available()False、两台机器结果不一样——三件事的第一嫌疑都是环境:终端里的 python 与 IDE 里的不是同一个、装了 CPU 版的 torch、依赖版本没锁。pip freeze > requirements.lock 把实验时的版本存进 run 目录(第六篇”能复现”的一部分)。环境、打包与交付的完整做法在 Infra 01 第七篇。

三、流式过一遍语料

1. 两种读法

语料通常是 JSONL:一行一个 JSON 对象。最直接的读法是整个文件读成一个 list

records = [json.loads(line) for line in path.read_text().splitlines()]   # 全部进内存

另一种是生成器:函数体里有 yield,调用它不执行,返回一个可以 for 的对象,每次 for 要下一个时才执行到下一个 yield

def iter_jsonl(path):
    with open(path, encoding="utf-8") as f:
        for line in f:                 # 文件对象本身就是逐行的迭代器,不会一次读完
            yield json.loads(line)

脚本合成了 10 万行、19.4 MB 的语料,两种读法各过一遍,用 tracemalloc 量峰值内存:

一次读进 list: 100,000 条, 峰值内存 102.9 MB  (≈ 文件大小 × 5.3)
生成器流式:   峰值内存 11.5 MB  (只有去重的哈希集合在涨)

list 的峰值是文件大小的 5 倍:每行变成一个 dict,每个键、每个字符串都是独立的 Python 对象,各带几十字节的头。10 GB 的语料按这个比例要 50 GB 内存。生成器一次只在内存里放一行,峰值与文件大小无关——那 11.5 MB 几乎全是去重用的哈希集合。

2. 串成流水线

生成器可以套生成器,每一层只管一件事:

def clean(records, min_words=5):
    seen = set()
    for r in records:
        if len(r["text"].split()) < min_words: continue          # 过滤
        h = hashlib.md5(r["text"].encode()).hexdigest()
        if h in seen: continue                                   # 精确去重
        seen.add(h); yield r

by_source, lengths = Counter(), Counter()
for r in clean(iter_jsonl(path)):                                # 读 → 过滤 → 去重 → 统计,一行一行流过去
    by_source[r["source"]] += 1
    lengths[min(len(r["text"].split()) // 10 * 10, 50)] += 1
读文件 ──行──▶ json.loads ──dict──▶ 过滤 ──▶ 去重 ──▶ Counter
          每个箭头上同一时刻只有一条记录;内存里常驻的只有 seen 和两个 Counter
过滤 + 去重后保留 92,867 / 100,000 条; 按来源: {'book': 30964, 'code': 31155, 'web': 30748}
按长度分桶(词数下界): 0+: 8339, 10+: 16465, 20+: 16754, 30+: 16550, 40+: 16529, 50+: 18230

这就是数据工程的最小形态:读、过滤、去重、统计四步,每步一个生成器。预训练系列第三篇的数据流水线——质量过滤、MinHash 近似去重(L2 第五篇)、配比——是同一个骨架换上更重的每一步。Counterdict 的子类,Counter()[k] += 1 不用先判断键在不在;pathlib.Pathpath / "train.jsonl"path.stat().st_sizepath.exists() 不用拼字符串。

3. 与 datasets 的关系

第五篇的 datasets 库把这一套做成了 load_dataset(..., streaming=True) + .filter() + .map():底层同样是逐条流过去(Arrow 格式、内存映射),接口是链式调用。知道它在做生成器的事,就知道为什么 streaming=True 的数据集没有 len()、为什么 .map() 不立即执行。

四、配置:dataclass

训练脚本的参数——模型名、学习率、batch、步数、LoRA 目标层——要能有默认值、能从命令行覆盖、能存进 run 目录、能被 IDE 补全@dataclass 一次满足:

from dataclasses import dataclass, field, asdict, replace

@dataclass
class TrainConfig:
    model: str = "Qwen/Qwen2.5-0.5B"
    lr: float = 2e-5
    batch_size: int = 8
    max_steps: int = 1000
    warmup: int = 50
    lora_targets: list[str] = field(default_factory=lambda: ["q_proj", "v_proj"])

    def __post_init__(self):
        assert self.warmup <= self.max_steps, f"warmup {self.warmup} > max_steps {self.max_steps}"

装饰器 @dataclass 读类体里的类型标注,自动生成 __init____repr____eq__。三个日常操作:

cfg  = TrainConfig()                                   # 默认
cfg2 = replace(cfg, lr=1e-4, max_steps=200)            # 覆盖:返回新对象,原来的不动
json.dump(asdict(cfg2), open(run_dir / "config.json", "w"))   # 存盘:dataclass → dict → JSON
覆盖后: 0.0001 200 | 原来的: 2e-05 1000
两个默认对象的 lora_targets 是否同一个 list: False
非法配置在构造时就报错: warmup 500 > max_steps 200

两个细节值得记住。可变默认值要用 field(default_factory=...):写成 lora_targets: list = ["q_proj"] 会让所有实例共享同一个 list,一个改了全改(dataclass 会直接拒绝这种写法并报错)。校验放在 __post_init__:非法组合在构造时就炸,而不是训了 500 步之后。transformersTrainingArgumentspeftLoraConfig 都是 dataclassTrainer(args=TrainingArguments(...)) 传的就是这样一个对象;命令行解析用 argparseHfArgumentParser(直接从 dataclass 的字段生成参数),配置文件用 JSON / YAML 读成 dictTrainConfig(**d) 展开。

类型标注 lr: float 在运行时不做检查(传字符串也能构造),它服务的是阅读、IDE 补全与 dataclass 这类读标注的工具;要做校验用 pydantic。类型系统的完整讨论在 Infra 01 第二篇。

五、训练代码里的六个语法

1. 对照表

PyTorch 的每个核心 API 都建在一个 Python 协议上。左边是你在训练代码里看到的,右边是它要求你写的:

flowchart LR
    subgraph P["PyTorch 里看到的"]
        A1["Dataset"]
        A2["DataLoader<br/>for batch in loader"]
        A3["nn.Module<br/>model(x)"]
        A4["@torch.no_grad()<br/>@torch.compile"]
        A5["with torch.autocast(...)<br/>with torch.no_grad()"]
        A6["Trainer(**kwargs)<br/>model.generate(**inputs)"]
    end
    subgraph Y["Python 要求你写 / 你要会读的"]
        B1["__len__ + __getitem__<br/>(序列协议)"]
        B2["__iter__ / 生成器<br/>(迭代协议)"]
        B3["__call__ → forward<br/>(可调用对象)"]
        B4["装饰器<br/>fn = deco(fn)"]
        B5["上下文管理器<br/>__enter__ / __exit__"]
        B6["*args / **kwargs<br/>(参数打包与展开)"]
    end
    A1 --> B1
    A2 --> B2
    A3 --> B3
    A4 --> B4
    A5 --> B5
    A6 --> B6
    classDef torch fill:#fff7e0,stroke:#c98a00,color:#222
    classDef py fill:#f4f8ff,stroke:#5b8def,color:#222
    class A1,A2,A3,A4,A5,A6 torch
    class B1,B2,B3,B4,B5,B6 py

2. 一个 40 行的”玩具 PyTorch”

脚本用纯 Python 把左列每一样各写了一个最小版,跑起来与真的形状一致:

class ToyDataset:                                   # Dataset:两个方法就够
    def __init__(self, texts): self.texts = texts
    def __len__(self): return len(self.texts)
    def __getitem__(self, i): return {"input_ids": [ord(c) % 128 for c in self.texts[i]], "label": len(self.texts[i]) % 2}

def loader(ds, batch_size, shuffle, seed=0):        # DataLoader 的骨架:一个生成器
    idx = list(range(len(ds)))
    if shuffle: random.Random(seed).shuffle(idx)
    for s in range(0, len(idx), batch_size):
        yield collate([ds[i] for i in idx[s:s + batch_size]])       # collate:一批样本拼成 batch,pad 到最长

class ToyModel:                                     # nn.Module:model(x) 走 __call__,__call__ 再调 forward
    def __call__(self, batch):
        self.calls += 1                             # 真实的 __call__ 在这里跑 forward hooks
        return self.forward(batch)
    def forward(self, batch): ...

def timed(fn):                                      # 装饰器:@timed 等价于 one_epoch = timed(one_epoch)
    @wraps(fn)
    def wrapper(*args, **kwargs):                   # *args / **kwargs:原样接住任何参数再原样传下去
        t0 = time.perf_counter(); out = fn(*args, **kwargs)
        print(f"[{fn.__name__} 用时 {time.perf_counter() - t0:.3f}s]"); return out
    return wrapper

@contextmanager
def seeded(seed):                                   # 上下文管理器:进入时做一件事,退出时(哪怕出错)恢复
    state = random.getstate(); random.seed(seed)
    try: yield
    finally: random.setstate(state)
len(ds) = 9; ds[0] = {'input_ids': [97, 116, 116, ...], 'label': 1}
第一个 batch: input_ids 形状 [4, 25], labels = [1, 0, 1, 0]
一个 epoch 看了 9 个样本, model 被调用 3 次 (= ceil(9 / 4))
seeded(42) 两次得到相同的数: True; 退出后随机状态已恢复

逐个读:

  • __len__ / __getitem__:写了这两个方法,len(ds)ds[3]for x in ds 就都能用——Python 见到 ds[3] 就调 ds.__getitem__(3)torch.utils.data.Dataset 要的就是这两个,DataLoader 按索引来取。
  • 生成器 / 迭代协议for batch in loader 每次要下一个 batch 时才取样本、才 collate——所以 DataLoader 不会把整个数据集拼好放内存里;IterableDataset 就是让你自己写 __iter__(一个生成器),第三章的流式读取直接能当它用。
  • __call__model(x)model.__call__(x)nn.Module__call__ 里先跑 hooks 再调你写的 forward。这就是为什么永远写 model(x) 而不是 model.forward(x)——后者跳过了 hooks(register_forward_hooktorch.compile 的一部分机制都挂在那里)。
  • 装饰器@timed 是”函数包函数”的语法糖,@torch.no_grad() 是同一个形状——返回一个进入时关梯度、退出时恢复的包装函数。@dataclass@torch.compile@functools.lru_cache@app.route 全是它。
  • 上下文管理器with torch.no_grad():with torch.autocast(...):with open(...)——进入时改一个状态,退出时保证恢复,中间抛异常也恢复。@contextmanager 把一个 yield 前后各一段的生成器变成它。
  • *args / **kwargs*args 把多余的位置参数收成 tuple,**kwargs 把多余的关键字参数收成 dict;调用时 f(*t, **d) 反过来展开。Trainer(**config)model.generate(**inputs)tokenizer(text, **kw) 都是把一个 dict 原样透传下去——看到它就去找那个 dict 里有什么键。

六、多进程预处理

1. 数字

tokenize、正则清洗、哈希这类CPU 密集的预处理,单进程跑 10 万行 0.46 秒,一亿行就是 8 分钟。脚本把同一个函数用三种方式跑:

total = sum(map(tokenize_count, lines))                        # 串行
with ThreadPool(8) as pool: pool.map(tokenize_count, lines, chunksize=2000)   # 8 线程
with Pool(8) as pool:       pool.map(tokenize_count, lines, chunksize=2000)   # 8 进程
100,000 行, 共 3,146,454 个 token; 8 个 worker
串行:   0.46s
线程池: 0.46s  (1.0×,GIL 让 CPU 密集的线程几乎不并行)
进程池: 0.15s  (3.1×,进程各有一个解释器,代价是启动与序列化)

2. 为什么

  • 线程 1.0×:CPython 有一把全局解释器锁(GIL),同一时刻只有一个线程在执行 Python 字节码。线程对等待(网络、磁盘、等 GPU)有用,对没用。这也是 DataLoader(num_workers=4) 开的是进程而不是线程的原因。
  • 进程 3.1× 而不是 8×:每个进程是一个独立解释器,要启动、要把输入 pickle 过去、把结果 pickle 回来。任务越轻,序列化占比越大;chunksize 让每次传一批而不是一条,是最重要的调节旋钮。任务重(每条几毫秒以上)时能接近核数。
  • 进程间不共享内存:全局变量在子进程里是拷贝,改了主进程看不见;传给 Pool.map 的函数必须是模块顶层可导入的(lambda 不行,因为要 pickle 函数本身)。

datasets.map(num_proc=8) 就是这个 Pool;GIL 的来历、asyncio 在什么时候比线程更合适、进程池与 DataLoader worker 的内部,在 Infra 01 第三篇。

七、出错的时候

1. traceback 从下往上读

Traceback (most recent call last):
  File ".../00_python_in_use.py", line 286, in exp_traceback
    train_step([[1.0, 2.0, 3.0], [4.0, 5.0]])
  File ".../00_python_in_use.py", line 279, in train_step
    return forward(batch)
  File ".../00_python_in_use.py", line 274, in forward
    h = project(batch, [[0.1] * 4 for _ in range(3)])   # w: [3, 4]
  File ".../00_python_in_use.py", line 269, in project
    assert len(row) == d_in, f"shape mismatch: ..."
AssertionError: shape mismatch: x row has 2 features, w expects 3

最后一行是错误本身(哪一类、什么信息),往上第一帧是出错的位置,再往上是它怎么被一层层调到的。PyTorch 的 traceback 常有二三十层,中间大半在 torch/nn/modules/module.py_call_impl 里——那是 __call__forward 的机制代码,跳过;找你自己文件出现的最后一帧。

2. 三类最常见的错

长什么样 第一反应
形状 mat1 and mat2 shapes cannot be multiplied (32x768 and 1024x768) 在出错前一行 print(x.shape, w.shape);第二篇的形状规则
设备 Expected all tensors to be on the same device, but found cuda:0 and cpu 某个张量忘了 .to(device)——常见于手建的 mask 或 label
类型 expected scalar type Float but found BFloat16 autocast 之外把 bf16 与 fp32 混算了;第四篇

写代码时在关键处 assert x.shape == (B, T, d), x.shape:形状错误在 PyTorch 里经常不报错(第二篇”能跑但错”),assert 让它在第一时间炸。需要停下来看变量时,在那一行前写 breakpoint(),运行到那里进入调试器(p x.shapen 下一行、c 继续)。测试与调试的系统做法在 Infra 01 第六篇。

八、越过哪条线进 Infra 01

本篇讲”怎么用”,每一节背后都有一个”为什么是这样”,那是 Infra 地图 01 系列的内容——两张地图共享,紧接本系列发布:

本篇 你会用了 想知道机制去
生成器、迭代协议、with open 01 第一篇:生成器怎么暂停恢复、迭代协议、名字绑定
dataclass、类型标注 01 第二篇:类型系统、Protocol、pydantic 与数据契约
Pool、GIL、chunksize 01 第三篇:GIL、线程 / 进程 / asyncio 的选择、DataLoader worker
装饰器、__call____getitem__ 01 第四篇:描述符、元类、算子注册表怎么用装饰器实现
峰值内存、对象开销 01 第五篇:引用计数、对象头、为什么一个 dict 比它的 JSON 大 5 倍
traceback、assertbreakpoint() 01 第六篇:pytest、性能剖析、线上排障
venv、requirements 01 第七篇:打包、pyproject、镜像与交付

算法工作的日常在左边两列就够了;读框架源码、给框架提 PR、排查 DataLoader 卡死这类问题时,右边那一列是必需的。

九、本文小结

  • 环境:一个项目一个环境,python -m pipimport torch 出问题先怀疑环境。
  • 流式过语料:一次读进 list 的峰值内存是文件大小的 5 倍(19 MB → 103 MB),生成器与文件大小无关(12 MB);读 → 过滤 → 去重 → 统计每步一个生成器串起来,是数据工程的最小形态。
  • 配置@dataclass 读类型标注生成 __init__ / __repr__replace 覆盖、asdict 存盘、__post_init__ 校验;可变默认值用 field(default_factory=...)
  • 六个语法__len__ / __getitem__Dataset;生成器是 DataLoader__call__forwardnn.Module(所以写 model(x));装饰器是 @torch.no_grad();上下文管理器是 with autocast**kwargs 是配置透传。
  • 多进程:GIL 让 CPU 密集的线程不并行(1.0×);进程池 8 个 worker 3.1×,差在启动与 picklechunksize 是旋钮;DataLoader(num_workers)datasets.map(num_proc) 都是它。
  • 出错:traceback 最后一行是错、往上第一帧是位置、找自己文件的最后一帧;形状 / 设备 / 类型三类错各有第一反应;assert 形状,breakpoint() 停下来看。

十、自测

  1. 一个 30 GB 的 JSONL 文件,机器内存 32 GB,json.loads 每行后要统计各来源的条数。能不能做?用什么写法?

    答案

    能。用生成器逐行读(文件对象本身就是逐行迭代器),Counter 累加——峰值内存只有一行加一个 Counter,与 30 GB 无关。读成 list 需要约 150 GB(5 倍),不行。

  2. for batch in DataLoader(ds, batch_size=8) 时,ds.__getitem__ 会在什么时候被调用?一共调多少次?

    答案

    for 每次要下一个 batch 时才被调用(迭代是惰性的),每个 batch 调 8 次,一个 epoch 共 len(ds) 次。数据集不会被提前全部取出来放内存。

  3. model.forward(x)model(x) 的输出一样,为什么仍然要写后者?

    答案

    model(x)nn.Module.__call__,在调 forward 前后执行 forward hooks(register_forward_hook、一些 profiler 与 torch.compile 的机制挂在那里);直接调 forward 跳过了它们,行为在挂了 hook 时会不同。

  4. 把一个 CPU 密集的清洗函数从 map 改成 ThreadPool(8).map,速度几乎不变。为什么?改什么能变?

    答案

    GIL:同一时刻只有一个线程执行 Python 字节码,CPU 密集的线程不并行。改用 multiprocessing.Pool(每个进程独立解释器),并设合适的 chunksize 减少序列化开销;任务越重越接近核数倍。

  5. @dataclass class C: tags: list = [] 会怎样?正确写法是什么?

    答案

    dataclass 直接拒绝并抛 ValueError: mutable default ... use default_factory——因为默认值只创建一次,所有实例会共享同一个 list。正确写法 tags: list = field(default_factory=list)

下一篇进入科学计算栈:在 NumPy 上建立形状直觉——轴、广播、einsum,手写一个 causal attention 并与 PyTorch 对数值;然后用 Pandas 分析评测结果、用 Matplotlib 看训练曲线。

  1. 它们各是一个 Python 协议,PyTorch 建在上面:__len__ / __getitem__Dataset 的全部要求,DataLoader 按索引来取;yield 定义生成器,for batch in loader 每次要下一个才算下一个,所以数据不会一次全进内存;@torch.no_grad() 是装饰器——”函数包函数”,进入时关梯度、退出时恢复;with autocast(...) 是上下文管理器,进入改状态、退出(含异常)保证恢复;**kwargs 把一个 dict 原样透传给下一层,看到它就去找那个 dict 里有什么键。另外 model(x)__call__ 再到 forward,hooks 挂在中间,所以不要直接调 forward。详见第五章。 

  2. 用生成器逐行读,读 → 过滤 → 去重 → 统计每步一个生成器串起来,同一时刻内存里只有一条记录加去重用的哈希集合。实测 19.4 MB 的 JSONL 读成 list 峰值 102.9 MB(文件大小的 5.3 倍,每个 dict、每个字符串都是带头的 Python 对象),生成器 11.5 MB 且与文件大小无关;按 5 倍算,10 GB 读成 list 要 50 GB,流式几十 MB 就够。精确去重靠内容哈希的集合;近似去重(MinHash)在 L2 第五篇。详见第三章。 

  3. 线程没用是因为 GIL:CPython 同一时刻只让一个线程执行字节码,CPU 密集的任务 8 线程 1.0×;线程只对等待(I/O、等 GPU)有用。进程池每个 worker 是独立解释器,实测 8 个 worker 3.1×,差在进程启动、输入输出的 pickle 序列化——任务越轻占比越大;chunksize 让每次传一批,任务越重越接近核数倍。DataLoader(num_workers)datasets.map(num_proc) 用的都是进程池。详见第六章。 

本文由 arganzheng 创作,采用 CC BY 4.0 许可协议。在保留原文作者、署名以及完整原文链接(https://arganzheng.life/python-in-use-for-algorithm-engineers.html)的前提下,欢迎各种形式的转载、翻译或商业引用。


COMMENTS

评论存放在 GitHub Discussions, 用 GitHub 账号登录即可发表,支持 Markdown。 想针对正文某句话说?选中那段文字,点浮出的「评论」即可划线评论;觉得哪里写错了,发表时勾上「同时提交 Issue」。 有人回复你时 GitHub 会按你的通知设置发邮件,不用守在这里。

×