使用Orbax恢复Flax模型检查点时触发ValueError问题排查
GPU训练的Orbax Checkpoint在CPU加载时的报错问题及解决
问题背景
使用以下代码在GPU训练期间保存模型状态,GPU环境加载正常,但切换到CPU环境加载时触发报错:
from flax.training import orbax_utils import orbax.checkpoint directory_gen_path = "checkpoints_loc" orbax_checkpointer_gen = orbax.checkpoint.PyTreeCheckpointer() gen_options = orbax.checkpoint.CheckpointManagerOptions(save_interval_steps=5, create=True) gen_checkpoint_manager = orbax.checkpoint.CheckpointManager( directory_gen_path, orbax_checkpointer_gen, gen_options ) def save_model_checkpoints(step_, generator_state, generator_batch_stats): gen_ckpt = { "model": generator_state, "batch_stats": generator_batch_stats, } save_args_gen = orbax_utils.save_args_from_target(gen_ckpt) gen_checkpoint_manager.save(step_, gen_ckpt, save_kwargs={"save_args": save_args_gen}) def load_model_checkpoints(generator_state, generator_batch_stats): gen_target = { "model": generator_state, "batch_stats": generator_batch_stats, } latest_step = gen_checkpoint_manager.latest_step() gen_ckpt = gen_checkpoint_manager.restore(latest_step, items=gen_target) generator_state = gen_ckpt["model"] generator_batch_stats = gen_ckpt["batch_stats"] return generator_state, generator_batch_stats
报错信息
- 旧版orbax-checkpoint:
ValueError: SingleDeviceSharding with Device=cuda:0 was not found in jax.local_devices().
- 更新至0.8.0版本后:
ValueError: sharding passed to deserialization should be specified, concrete and an instance of
jax.sharding.Sharding. Got None
原因分析
- 核心问题是checkpoint保存了GPU设备的分片(Sharding)信息:训练时模型参数存储在
cuda:0设备,Orbax默认会将参数的分片配置一并保存。当在CPU环境加载时,系统找不到cuda:0设备,旧版直接抛出设备不存在的错误;新版则要求必须明确指定兼容的分片策略,不允许默认使用None。 - 加载时未指定CPU兼容的分片规则,导致Orbax尝试复用训练时的GPU分片配置,与当前环境不匹配。
解决思路
方案1:保存时忽略分片信息(推荐)
在保存checkpoint时,主动移除分片信息,让checkpoint不绑定特定设备,这样在任意环境都能加载。修改save_model_checkpoints函数:
def save_model_checkpoints(step_, generator_state, generator_batch_stats): gen_ckpt = { "model": generator_state, "batch_stats": generator_batch_stats, } # 生成save_args,指定不保存sharding信息 save_args_gen = orbax_utils.save_args_from_target( gen_ckpt, save_args_fn=lambda x: orbax.checkpoint.SaveArgs(sharding=None) ) gen_checkpoint_manager.save(step_, gen_ckpt, save_kwargs={"save_args": save_args_gen})
方案2:加载时指定CPU分片
如果无法重新生成checkpoint,可在CPU加载时明确指定CPU的分片策略,覆盖原checkpoint中的GPU分片配置:
def load_model_checkpoints(generator_state, generator_batch_stats): gen_target = { "model": generator_state, "batch_stats": generator_batch_stats, } # 获取CPU设备并创建分片 cpu_device = jax.devices('cpu')[0] cpu_sharding = jax.sharding.SingleDeviceSharding(cpu_device) # 为每个参数指定CPU分片 restore_args = orbax_utils.restore_args_from_target( gen_target, restore_args_fn=lambda x: orbax.checkpoint.RestoreArgs(sharding=cpu_sharding) ) latest_step = gen_checkpoint_manager.latest_step() gen_ckpt = gen_checkpoint_manager.restore( latest_step, items=gen_target, restore_kwargs={"restore_args": restore_args} ) generator_state = gen_ckpt["model"] generator_batch_stats = gen_ckpt["batch_stats"] return generator_state, generator_batch_stats
方案3:加载前将目标参数移至CPU
在调用restore前,先把传入的generator_state和generator_batch_stats移到CPU设备,让Orbax自动匹配目标设备的分片:
def load_model_checkpoints(generator_state, generator_batch_stats): # 将目标参数移至CPU generator_state = jax.device_get(generator_state) generator_batch_stats = jax.device_get(generator_batch_stats) gen_target = { "model": generator_state, "batch_stats": generator_batch_stats, } latest_step = gen_checkpoint_manager.latest_step() gen_ckpt = gen_checkpoint_manager.restore(latest_step, items=gen_target) generator_state = gen_ckpt["model"] generator_batch_stats = gen_ckpt["batch_stats"] return generator_state, generator_batch_stats
内容的提问来源于stack exchange,提问作者yash
相关产品推荐
相关产品推荐

