news 2026/9/10 16:12:12

Transformers 大模型实例化:低内存检查点分片(Sharding)与加载机制全解

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Transformers 大模型实例化:低内存检查点分片(Sharding)与加载机制全解

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. 加载(预训练好的)权重;
  3. 将这些预训练权重填入(放置到)随机模型中。

步骤 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_keysunexpected_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_keysunexpected_keysNamedTuple_IncompatibleKeys),便于调用方复用既有校验逻辑。

该功能在仓库测试中亦有覆盖,例如 tests/utils/test_modeling_utils.py 中大量用例使用model.save_pretrained(tmp_dir, max_shard_size="100kB")等方式验证分片保存与重载路径,可作为实际行为的参考依据。

低内存加载(Low Memory Loading)

分片检查点解决的是上文工作流步骤 2的内存占用问题。若希望在低内存环境中使用该模型(即连“模型大小 + 最大分片”的占用都想进一步压低,例如步骤 1 的随机模型创建也走低内存路径),官方指南建议借助基于Accelerate 库的加载工具,具体包括device_maplow_cpu_mem_usageload_in_8bitfrom_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_maplow_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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/10 16:11:06

Telegram Bot API自定义扩展终极指南:如何快速添加新的API端点和方法

Telegram Bot API自定义扩展终极指南:如何快速添加新的API端点和方法 Telegram Bot API是一个功能强大的机器人开发框架,而telegram-bot-api项目提供了Golang语言的完整绑定支持。📱 本文将详细介绍如何在这个库中自定义扩展,添加…

作者头像 李华
网站建设 2026/9/10 16:11:04

会议录音不发云端:Buzz 三步离线转文字的完整走法

会议录音不发云端:Buzz 三步离线转文字的完整走法 【免费下载链接】buzz Buzz transcribes and translates audio offline on your personal computer. Powered by OpenAIs Whisper. 项目地址: https://gitcode.com/GitHub_Trending/buz/buzz Buzz 是一款离线…

作者头像 李华
网站建设 2026/9/10 16:10:32

Calibre 格式转换完全指南:从单本书到整个书库

Calibre 格式转换完全指南:从单本书到整个书库 【免费下载链接】calibre The official source code repository for the calibre ebook manager 项目地址: https://gitcode.com/GitHub_Trending/ca/calibre 你和朋友共用一个书架文件夹。对方的 Kindle 只认 …

作者头像 李华
网站建设 2026/9/10 16:09:02

ATAC-seq技术解析与马铃薯耐寒研究应用

1. ATAC-seq技术解析:打开染色质可及性研究的钥匙ATAC-seq(Assay for Transposase-Accessible Chromatin using sequencing)是近年来表观遗传学领域的一项革命性技术。这项技术的核心在于利用改造后的Tn5转座酶,特异性切割开放染色…

作者头像 李华