Jax/Flax下Orbax checkpoint恢复报错,求正确恢复方法
Orbax Checkpoint恢复失败的解决方案
问题场景
用户使用以下代码保存Orbax Checkpoint:
check_options = ocp.CheckpointManagerOptions(max_to_keep=5, create=True) check_path = Path(os.getcwd(), out_dir, 'checkpoint') checkpoint_manager = ocp.CheckpointManager(check_path, options=check_options, item_names=('state', 'metadata')) checkpoint_manager.save( step=iter_num, args=ocp.args.Composite( state=ocp.args.StandardSave(state), metadata=ocp.args.JsonSave((model_args, iter_num, best_val_loss, losses['val'].item(), config))))
恢复时使用这段代码获取state:
state, lr_schedule = init_train_state(model, params['params'], learning_rate, weight_decay, beta1, beta2, decay_lr, warmup_iters, lr_decay_iters, min_lr) # 此处state是初始化后的Train_state类型变量 state = checkpoint_manager.restore(checkpoint_manager.latest_step(), items={'state': state})
但在训练循环中使用恢复后的state时,触发以下错误:
--------------------------------------------------------------------------- KeyError Traceback (most recent call last) File /opt/conda/envs/py_3.10/lib/python3.10/site-packages/jax/_src/api_util.py:584, in shaped_abstractify(x) 583 try: --> 584 return _shaped_abstractify_handlers[type(x)](x) 585 except KeyError: KeyError: <class 'orbax.checkpoint.composite_checkpoint_handler.CompositeArgs'> During handling of the above exception, another exception occurred: TypeError Traceback (most recent call last) Cell In[40], line 37 34 if iter_num == 0 and eval_only: 35 break ---> 37 state, loss = train_step(state, get_batch('train')) 39 # timing and logging 40 t1 = time.time() [... skipping hidden 6 frame] File /opt/conda/envs/py_3.10/lib/python3.10/site-packages/jax/_src/api_util.py:575, in _shaped_abstractify_slow(x) 573 dtype = dtypes.canonicalize_dtype(x.dtype, allow_extended_dtype=True) 574 else: --> 575 raise TypeError( 576 f"Cannot interpret value of type {type(x)} as an abstract array; it " 577 "does not have a dtype attribute") 578 return core.ShapedArray(np.shape(x), dtype, weak_type=weak_type, 579 named_shape=named_shape) TypeError: Cannot interpret value of type <class 'orbax.checkpoint.composite_checkpoint_handler.CompositeArgs'> as an abstract array; it does not have a dtype attribute
问题原因
调用checkpoint_manager.restore并指定items参数时,返回的是CompositeArgs对象,而非直接返回恢复后的state实例。直接将该对象赋值给state变量,会导致后续JAX训练逻辑无法识别这个非数组类型的对象,从而触发类型错误。
解决方法
需要从返回的CompositeArgs对象中提取实际的state值,通过属性访问的方式获取:
修正后的恢复代码:
state, lr_schedule = init_train_state(model, params['params'], learning_rate, weight_decay, beta1, beta2, decay_lr, warmup_iters, lr_decay_iters, min_lr) # 先获取恢复结果对象,再提取state restored_result = checkpoint_manager.restore(checkpoint_manager.latest_step(), items={'state': state}) state = restored_result.state
或者简化为一行:
state = checkpoint_manager.restore(checkpoint_manager.latest_step(), items={'state': state}).state
验证方法
恢复后可以打印state的类型,确认是你的Train_state类型而非CompositeArgs:
print(type(state)) # 输出应与初始化后的state类型一致
内容的提问来源于stack exchange,提问作者Dmitry J
相关产品推荐
相关产品推荐

