【Bug已解决】[Feature] Save model-only feature in `save_state` 解决方案

【Bug已解决】[Feature] Save model-only feature insave_state解决方案

一、现象长什么样

accelerator.save_state("ckpt")保存训练状态时,发现它把优化器、学习率调度器、RNG 状态全打包了

  • 一个 7B 模型(权重14GB fp16),save_state产出的目录却是 **40GB**——因为 Adam 优化器的 m/v 状态是参数的 2 倍,加上调度器/RNG,体积暴涨。
  • 保存/加载很慢(要序列化几十 GB),但你可能只是想「存个模型权重给下游推理/微调用」。
  • 恢复时也强制要求优化器状态齐全,否则报「缺 key」,但你根本不需要优化器。

特征:

  • 只在大模型 + 想只存模型时痛点明显;小模型无所谓。
  • save_state没有「只存模型」的开关,用户要么全存、要么自己另写model.save_pretrained
  • 想用 Accelerate 的统一接口却被迫绕过它,割裂。

本质:accelerator.save_state把「模型 + 优化器 + 调度器 + RNG」绑死成一个原子保存,没有提供「只存模型权重」的选项。对只想存模型(推理/分发/下游微调)的场景,优化器等状态是冗余负担。

二、背景

accelerator.save_state(path)的设计目标是「完整恢复训练」——所以它存四类状态:

  1. 模型权重model.safetensors/pytorch_model.bin)。
  2. 优化器状态(Adam 的exp_avg/exp_avg_sq,通常 = 2× 参数量)。
  3. 学习率调度器状态last_epoch等)。
  4. RNG 状态(CPU/GPU 随机种子,保证可复现)。

对「从头恢复训练」这四类都得要。但对很多场景只要模型

  • 推理分发:把训好的模型发给别人做推理,不需要优化器。
  • 下游微调:用阶段 checkpoint 作为底座做 SFT,只需权重。
  • 快速快照:训练中想频繁存个「当前模型长啥样」看效果,不想每次都写几十 GB。

此时优化器那 2× 参数就是纯浪费(体积 + 时间)。这个 Feature Request 就是要求save_state支持save_model_only=True,只序列化模型权重,跳过优化器/调度器/RNG。

一句话:save_state缺少「只存模型」选项,大模型下优化器状态造成体积与时间浪费。

三、根因(能力缺口分析)

把这个缺口当 bug 分析,根因是save_state把四类状态绑死、无「仅模型」开关、且缺轻量保存路径,三层:

第一层(主因):四类状态原子保存,无 model-only 开关。save_state内部依次序列化 model/optim/scheduler/rng,没有if model_only: skip others的分支。用户要只存模型只能绕过 Accelerate 自己调model.save_pretrained

第二层:优化器状态体积被忽视。实现没意识到「Adam 状态 = 2× 参数」对大模型是几十 GB 的负担,没提供「跳过它」以提速省空间的选项。

第三层:load 与 save 不对称。即便你手存了「只有模型」的目录,load_state预期的是「四件套齐全」,缺优化器就报错。缺一个「model-only 保存 / 加载」对称的能力,导致用户绕过后更难恢复。

一句话:save_state 无 model-only 分支、忽视优化器体积、load/save 不对称,导致大模型只存模型时被迫全量序列化。

四、最小可运行复现

下面用纯 Python 模拟「save_state 全量序列化 vs model-only 跳过优化器」的体积差异,不需要 GPU:

from dataclasses import dataclass from typing import Dict @dataclass class StateSizes: model_gb: float = 14.0 optim_gb: float = 28.0 # Adam m/v = 2x scheduler_gb: float = 0.001 rng_gb: float = 0.001 def save_state_buggy(sizes: StateSizes, model_only: bool = False) -> Dict: out = {"model": sizes.model_gb} if not model_only: # 错误:无 model-only 分支,永远存优化器 out["optim"] = sizes.optim_gb out["scheduler"] = sizes.scheduler_gb out["rng"] = sizes.rng_gb return out def total_gb(d: Dict) -> float: return sum(d.values()) def main(): s = StateSizes() full = save_state_buggy(s, model_only=False) only = save_state_buggy(s, model_only=True) print(f"全量 save_state: {total_gb(full):.1f} GB") print(f"model-only(理想): {total_gb(only):.1f} GB") if __name__ == "__main__": main()

跑出来全量 ~42GB、model-only ~14GB——直观展示了「优化器状态占大头、只存模型可省 2/3」。

五、解决方案(第一层:最小直接修复)

最省事的救火:直接调model.save_pretrained只存模型,绕开save_state

from accelerate import Accelerator accelerator = Accelerator() # ... 训练 ... # 只想存模型权重(推理/下游用):用模型自带的保存,不碰优化器 accelerator.unwrap_model(model).save_pretrained("ckpt-model-only") # 若用了 prepare,注意 unwrap 拿到原始模型

恢复模型(不含优化器)时:

from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("ckpt-model-only") # 只加载权重

这能立刻省掉优化器的几十 GB。缺点是绕开了 Accelerate 的统一save_state接口,且若之后想完整恢复训练得另外存优化器。

六、解决方案(第二层:结构性改进)

第一层是「绕开存模型」,第二层是「实现 Feature:给save_statemodel_only开关,save/load 对称」,从设计上消灭全量绑死:

from dataclasses import dataclass from typing import Dict, Optional @dataclass class SavePolicy: model_only: bool = False def collect(self, model, optimizer=None, scheduler=None, rng=None) -> Dict: out = {"model": model} if self.model_only: return out # 只存模型,跳过其余 # 完整保存 if optimizer is not None: out["optimizer"] = optimizer if scheduler is not None: out["scheduler"] = scheduler if rng is not None: out["rng"] = rng return out def save_state_safe(path, model, optimizer=None, scheduler=None, rng=None, model_only: bool = False): policy = SavePolicy(model_only=model_only) state = policy.collect(model, optimizer, scheduler, rng) # 序列化 state 到 path _serialize(path, state) return state # 用法 save_state_safe("ckpt", model, optimizer, scheduler, rng, model_only=True) # 只存模型,~14GB 而非 ~42GB

配套load_state也支持model_only

def load_state_safe(path, model, optimizer=None, scheduler=None, rng=None, model_only: bool = False): state = _deserialize(path) model.load_state_dict(state["model"]) if not model_only: if "optimizer" in state and optimizer is not None: optimizer.load_state_dict(state["optimizer"]) # ... 其余 # model_only 时静默跳过优化器/调度器/RNG,不报错

关键改动:

  1. SavePolicy.model_only分支:只 collect 模型,跳过优化器/调度器/RNG。
  2. load_state对称支持model_only:缺优化器不报错,只加载模型。
  3. save/load 都用同一开关,能力对称、不割裂。

七、解决方案(第三层:断言 / CI 守护)

把「model-only 跳过优化器」「体积更小」「load 对称不报错」固化成测试:

import pytest def test_model_only_skips_optim(): policy = SavePolicy(model_only=True) state = policy.collect(model="M", optimizer="O", scheduler="S", rng="R") assert "optim" not in state assert "scheduler" not in state assert "rng" not in state assert state["model"] == "M" def test_full_includes_all(): policy = SavePolicy(model_only=False) state = policy.collect(model="M", optimizer="O", scheduler="S", rng="R") assert set(state) == {"model", "optim", "scheduler", "rng"} def test_model_only_smaller(): s = StateSizes() full = save_state_buggy(s, False) only = save_state_buggy(s, True) assert total_gb(only) < total_gb(full) def test_load_model_only_no_optim_ok(): # model_only 保存后,load 不要求优化器,不报错 state = {"model": "M"} # model-only 产物 loaded = {} if "model" in state: loaded["model"] = state["model"] # 优化器缺失但 model_only=True -> 不报错 assert "model" in loaded def test_save_load_symmetric(): # 同开关 save/load 一致 policy = SavePolicy(model_only=True) out = policy.collect("M", "O") assert "optim" not in out

再加一个端到端回归:model-only 保存体积远小于全量,且能独立加载模型:

def test_model_only_save_and_load(): save_state_safe("ckpt", model, optimizer, scheduler, rng, model_only=True) # 只加载模型成功(不要求优化器) model2 = load_model_only("ckpt") assert model2 is not None

八、排查清单

  1. save_state产出目录异常大(≈3× 模型权重)且你只需模型 → 是优化器状态冗余。
  2. 临时救火:用accelerator.unwrap_model(model).save_pretrained(...)只存模型。
  3. 长期方向:给save_statemodel_only=True开关,save/load 对称。
  4. 确认load_state在 model_only 下不因缺优化器报错。
  5. 升级 accelerate 到合了该开关的版本,并跑上面的test_model_only_skips_optim
  6. 若既要模型又要偶尔完整恢复,可「model_only 频繁存 + 定期全量存」混合策略。
  7. 注意 FSDP/TP 下模型是分片的,model-only 保存也要按分片保存(见之前分片议题)。

九、小结

save_state缺 model-only 选项,不是接口坏了,而是它把模型/优化器/调度器/RNG 绑死成原子保存,无「只存模型」开关,大模型下优化器状态(2× 参数)造成体积与时间浪费。最小修复是用unwrap_model(model).save_pretrained(...)只存模型;结构性方向是实现 Feature——给save_statemodel_only开关且 load/save 对称;最后用 pytest 把「model-only 跳过优化器」「体积更小」「load 对称不报错」锁死。抓住「保存粒度应与用途匹配、save/load 必须对称」这条,所有 checkpoint 体积/速度痛点都能照此优化。