公司动态
【Bug已解决】[Feature] Save model-only feature in `save_state` 解决方案
【Bug已解决】[Feature] Save model-only feature insave_state解决方案一、现象长什么样用accelerator.save_state(ckpt)保存训练状态时发现它把优化器、学习率调度器、RNG 状态全打包了一个 7B 模型权重14GB fp16save_state产出的目录却是 **40GB**——因为 Adam 优化器的 m/v 状态是参数的 2 倍加上调度器/RNG体积暴涨。保存/加载很慢要序列化几十 GB但你可能只是想「存个模型权重给下游推理/微调用」。恢复时也强制要求优化器状态齐全否则报「缺 key」但你根本不需要优化器。特征只在大模型 想只存模型时痛点明显小模型无所谓。save_state没有「只存模型」的开关用户要么全存、要么自己另写model.save_pretrained。想用 Accelerate 的统一接口却被迫绕过它割裂。本质accelerator.save_state把「模型 优化器 调度器 RNG」绑死成一个原子保存没有提供「只存模型权重」的选项。对只想存模型推理/分发/下游微调的场景优化器等状态是冗余负担。二、背景accelerator.save_state(path)的设计目标是「完整恢复训练」——所以它存四类状态模型权重model.safetensors/pytorch_model.bin。优化器状态Adam 的exp_avg/exp_avg_sq通常 2× 参数量。学习率调度器状态last_epoch等。RNG 状态CPU/GPU 随机种子保证可复现。对「从头恢复训练」这四类都得要。但对很多场景只要模型推理分发把训好的模型发给别人做推理不需要优化器。下游微调用阶段 checkpoint 作为底座做 SFT只需权重。快速快照训练中想频繁存个「当前模型长啥样」看效果不想每次都写几十 GB。此时优化器那 2× 参数就是纯浪费体积 时间。这个 Feature Request 就是要求save_state支持save_model_onlyTrue只序列化模型权重跳过优化器/调度器/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 跳过优化器」的体积差异不需要 GPUfrom 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_onlyFalse) only save_state_buggy(s, model_onlyTrue) print(f全量 save_state: {total_gb(full):.1f} GB) print(fmodel-only(理想): {total_gb(only):.1f} GB) if __name__ __main__: main()跑出来全量 ~42GB、model-only ~14GB——直观展示了「优化器状态占大头、只存模型可省 2/3」。五、解决方案第一层最小直接修复最省事的救火直接调model.save_pretrained只存模型绕开save_statefrom 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_state加model_only开关save/load 对称」从设计上消灭全量绑死from dataclasses import dataclass from typing import Dict, Optional dataclass class SavePolicy: model_only: bool False def collect(self, model, optimizerNone, schedulerNone, rngNone) - 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, optimizerNone, schedulerNone, rngNone, model_only: bool False): policy SavePolicy(model_onlymodel_only) state policy.collect(model, optimizer, scheduler, rng) # 序列化 state 到 path _serialize(path, state) return state # 用法 save_state_safe(ckpt, model, optimizer, scheduler, rng, model_onlyTrue) # 只存模型~14GB 而非 ~42GB配套load_state也支持model_onlydef load_state_safe(path, model, optimizerNone, schedulerNone, rngNone, 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不报错关键改动SavePolicy.model_only分支只 collect 模型跳过优化器/调度器/RNG。load_state对称支持model_only缺优化器不报错只加载模型。save/load 都用同一开关能力对称、不割裂。七、解决方案第三层断言 / CI 守护把「model-only 跳过优化器」「体积更小」「load 对称不报错」固化成测试import pytest def test_model_only_skips_optim(): policy SavePolicy(model_onlyTrue) state policy.collect(modelM, optimizerO, schedulerS, rngR) 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_onlyFalse) state policy.collect(modelM, optimizerO, schedulerS, rngR) 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_onlyTrue - 不报错 assert model in loaded def test_save_load_symmetric(): # 同开关 save/load 一致 policy SavePolicy(model_onlyTrue) 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_onlyTrue) # 只加载模型成功不要求优化器 model2 load_model_only(ckpt) assert model2 is not None八、排查清单看save_state产出目录异常大≈3× 模型权重且你只需模型 → 是优化器状态冗余。临时救火用accelerator.unwrap_model(model).save_pretrained(...)只存模型。长期方向给save_state加model_onlyTrue开关save/load 对称。确认load_state在 model_only 下不因缺优化器报错。升级 accelerate 到合了该开关的版本并跑上面的test_model_only_skips_optim。若既要模型又要偶尔完整恢复可「model_only 频繁存 定期全量存」混合策略。注意 FSDP/TP 下模型是分片的model-only 保存也要按分片保存见之前分片议题。九、小结save_state缺 model-only 选项不是接口坏了而是它把模型/优化器/调度器/RNG 绑死成原子保存无「只存模型」开关大模型下优化器状态2× 参数造成体积与时间浪费。最小修复是用unwrap_model(model).save_pretrained(...)只存模型结构性方向是实现 Feature——给save_state加model_only开关且 load/save 对称最后用 pytest 把「model-only 跳过优化器」「体积更小」「load 对称不报错」锁死。抓住「保存粒度应与用途匹配、save/load 必须对称」这条所有 checkpoint 体积/速度痛点都能照此优化。