You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

原因分析

  1. 核心问题是checkpoint保存了GPU设备的分片(Sharding)信息:训练时模型参数存储在cuda:0设备,Orbax默认会将参数的分片配置一并保存。当在CPU环境加载时,系统找不到cuda:0设备,旧版直接抛出设备不存在的错误;新版则要求必须明确指定兼容的分片策略,不允许默认使用None。
  2. 加载时未指定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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.16 11:26:16