Transformers 大模型实例化:低内存检查点分片(Sharding)与加载机制全解
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
本文基于 Transformers 官方指南 big_models.md 展开,系统讲解在实例化大型预训练模型时如何最小化 RAM 占用:包括检查点自动分片(sharded checkpoints)、max_shard_size参数控制、索引文件结构、load_sharded_checkpoint的逐片加载实现,以及基于 Accelerate 的更低内存加载方案,并结合当前仓库源码印证每一处机制的实际行为。
为什么加载大模型会“爆”内存
使用大型预训练模型时,控制 RAM 用量始终是核心难题。常规的 PyTorch 工作流分为三步:
- 创建一个带随机权重的模型;
- 加载(预训练好的)权重;
- 将这些预训练权重填入(放置到)随机模型中。
步骤 1 和步骤 2 都需要在内存中保存一份完整的模型副本。对于小模型这没有问题,但当模型体积达到数 GB 时,两份副本就可能超出 RAM 上限。更糟糕的是,如果使用torch.distributed进行分布式训练,每个进程都会各自加载一次预训练模型,即每个进程都要保存这两份副本,内存压力被进程数成倍放大。
官方指南中特别提示了随机初始化的一个细节:随机创建的模型在内存中实际上是用“空(empty)”张量初始化的,所谓“随机值”只是恰好读取了内存对应区块中已存在的数据。因此在步骤 3 中,如果模型/参数本身带有合适的初始化分布(如正态分布),未初始化权重的填充可以非常快。
Transformers 针对上述问题提供了一套分片加载解决方案。需要注意的是,该领域仍在持续演进中,未来相关 API 可能略有调整。
分片检查点(Sharded Checkpoints)
自 4.18.0 版本起,超过 10GB 的模型检查点在保存时会被自动拆分为多个较小的部分:调用model.save_pretrained(save_dir)时,除了单个完整检查点的情况,Transformers 会生成若干部分检查点(每个小于指定大小)以及一个将参数名映射到其所在文件的索引文件。当前仓库源码中save_pretrained的签名为 src/transformers/modeling_utils.py#L3212-L3224,max_shard_size参数当前默认值为"50GB",可通过max_shard_size参数控制分片前的单个检查点最大尺寸:
源码提示(src/transformers/modeling_utils.py#L3245-L3254):如果模型中存在单个权重就大于
max_shard_size的情况,该权重会被单独放入一个分片,此时该分片会大于max_shard_size。
用 BERT 验证分片效果
官方指南以传统 BERT 模型为例演示了分片保存。先加载模型:
from transformers import AutoModel model = AutoModel.from_pretrained("google-bert/bert-base-cased")使用 [~PreTrainedModel.save_pretrained] 默认保存时,生成的目录包含两个文件:模型配置信息和权重信息。
>>> import os >>> import tempfile >>> with tempfile.TemporaryDirectory() as tmp_dir: ... model.save_pretrained(tmp_dir) ... print(sorted(os.listdir(tmp_dir))) ['config.json', 'pytorch_model.bin']将最大分片大小设置为 200MB 后,结果变为 3 个权重分片加 1 个索引文件:
>>> with tempfile.TemporaryDirectory() as tmp_dir: ... model.save_pretrained(tmp_dir, max_shard_size="200MB") ... print(sorted(os.listdir(tmp_dir))) ['config.json', 'pytorch_model-00001-of-00003.bin', 'pytorch_model-00002-of-00003.bin', 'pytorch_model-00003-of-00003.bin', 'pytorch_model.bin.index.json']在模型配置之上,你会看到 3 个不同的权重文件与index.json索引文件。这样保存的分片检查点可以用 [~PreTrainedModel.from_pretrained] 方法完整重新加载:
>>> with tempfile.TemporaryDirectory() as tmp_dir: ... model.save_pretrained(tmp_dir, max_shard_size="200MB") ... new_model = AutoModel.from_pretrained(tmp_dir)核心收益:对大型模型而言,上述工作流的步骤 2 中,每个检查点分片在加载完前一个分片后才加载下一个,RAM 内存占用被限制在“模型大小 + 最大分片大小”的水平,而非两份完整模型副本。
索引文件的内部结构
内部实现依赖索引文件来决定哪个键存在于哪个检查点中、对应权重存储在哪个文件。该索引与普通 JSON 文件无异,可直接读取为字典:
>>> import json >>> with tempfile.TemporaryDirectory() as tmp_dir: ... model.save_pretrained(tmp_dir, max_shard_size="200MB") ... with open(os.path.join(tmp_dir, "pytorch_model.bin.index.json"), "r") as f: ... index = json.load(f) >>> print(index.keys()) dict_keys(['metadata', 'weight_map'])索引包含两个顶层键:
metadata:目前仅包含模型总大小(未来计划补充其他信息):
>>> index["metadata"] {'total_size': 433245184}weight_map:索引的主体,将每个参数名(即 PyTorch 模型state_dict中常见的键)映射到其所在文件:
>>> index["weight_map"] {'embeddings.LayerNorm.bias': 'pytorch_model-00001-of-00003.bin', 'embeddings.LayerNorm.weight': 'pytorch_model-00001-of-00003.bin', ... }手动加载分片检查点:load_sharded_checkpoint
如果不想在模型内部使用 [~PreTrainedModel.from_pretrained],而是像对完整检查点使用model.load_state_dict()那样直接加载分片检查点,应使用 [~trainer_utils.load_sharded_checkpoint]:
>>> from transformers.trainer_utils import load_sharded_checkpoint >>> with tempfile.TemporaryDirectory() as tmp_dir: ... model.save_pretrained(tmp_dir, max_shard_size="200MB") ... load_sharded_checkpoint(model, tmp_dir)从源码结构看(src/transformers/trainer_utils.py#L1057-L1130),该函数的行为与文档描述完全一致,几个值得注意的实现细节:
- 参数:
model为待加载的模型,folder为分片检查点所在目录,strict=True表示严格校验键匹配(在加载任何状态字典之前即报错,错误信息中会列出missing_keys与unexpected_keys),prefer_safe=True表示当目录中同时存在 safetensors 与 PyTorch 格式时优先加载 safetensors; - 逐片加载与内存释放(src/transformers/trainer_utils.py#L1121-L1127):按索引中的分片文件列表逐个加载,每个分片加载进模型后立即
del state_dict并执行gc.collect(),确保内存在下一次加载前被释放——这正是“RAM 占用 = 模型 + 最大分片”的底层保证; - 返回类型:与 PyTorch 原生
load_state_dict相同,返回含missing_keys和unexpected_keys的NamedTuple(_IncompatibleKeys),便于调用方复用既有校验逻辑。
该功能在仓库测试中亦有覆盖,例如 tests/utils/test_modeling_utils.py 中大量用例使用model.save_pretrained(tmp_dir, max_shard_size="100kB")等方式验证分片保存与重载路径,可作为实际行为的参考依据。
低内存加载(Low Memory Loading)
分片检查点解决的是上文工作流步骤 2的内存占用问题。若希望在低内存环境中使用该模型(即连“模型大小 + 最大分片”的占用都想进一步压低,例如步骤 1 的随机模型创建也走低内存路径),官方指南建议借助基于Accelerate 库的加载工具,具体包括device_map、low_cpu_mem_usage、load_in_8bit等from_pretrained参数。
详细用法请参考当前仓库文档中的使用 Accelerate 加载大模型章节(原文档相对链接./main_classes/model#large-model-loading,位于docs/source/ja/main_classes/model.md)。
小结:内存优化路径速查
| 阶段 | 问题 | 解决方案 | 关键入口 |
|---|---|---|---|
| 保存 | 单文件检查点过大,难以在分布式/网络环境传输 | max_shard_size自动分片 +index.json | ~PreTrainedModel.save_pretrained |
| 加载(步骤 2) | 完整权重文件需一次性进 RAM | 逐分片顺序加载,RAM 占用限制为模型 + 最大分片 | [~PreTrainedModel.from_pretrained] /~trainer_utils.load_sharded_checkpoint |
| 加载(步骤 1+3) | 随机模型 + 权重两份完整副本 | 基于 Accelerate 的低内存加载(device_map、low_cpu_mem_usage等) | 模型加载指南 |
适用前提与限制:分片加载依赖检查点目录中存在pytorch_model.bin.index.json(或 safetensors 对应的model.safetensors.index.json),缺少索引文件时load_sharded_checkpoint会直接抛出ValueError;对单个权重即超过max_shard_size的极端情况,分片不会强行拆分该权重。由于该 API 仍在演进,接入新版本 Transformers 时建议对照仓库内 src/transformers/trainer_utils.py 与 src/transformers/modeling_utils.py 的最新签名确认参数行为。
【免费下载链接】transformers🤗 Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考