mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-28 20:11:22 +00:00
https://github.com/microsoft/onnxruntime/pull/4639 changed the default behavior by removing optimizer state from state_dict/checkpoint APIs. The reason for the previous change was to allow models trained on ORT to be used for inference on PyTorch, which is an important feature. Due to the change aforementioned, when resuming training from a checkpoint, the optimizer would start with random weights, leading to a bad performance. This behavior would also cause reproducibility issues, as the optimizer wouldnt be able to resume from its previous state. This PR adds a boolean flag to state_dict/save_xheckpoint API that when True (default) it saves both model and optimizer state. When False, only the model state is kept.
This commit is contained in:
parent
8679a7244e
commit
9817b8c8a7
1 changed files with 9 additions and 10 deletions
|
|
@ -473,11 +473,11 @@ def create_ort_training_session_with_optimizer(model, device, training_optimizer
|
|||
|
||||
return session, train_io_binding, eval_io_binding, output_name, torch_params, output_types
|
||||
|
||||
def save_checkpoint(model, checkpoint_dir, checkpoint_prefix="ORT_checkpoint", checkpoint_state_dict=None):
|
||||
def save_checkpoint(model, checkpoint_dir, checkpoint_prefix="ORT_checkpoint", checkpoint_state_dict=None, include_optimizer_state=True):
|
||||
if checkpoint_state_dict==None:
|
||||
checkpoint_state_dict={'model': model.state_dict()}
|
||||
checkpoint_state_dict={'model': model.state_dict(include_optimizer_state)}
|
||||
else:
|
||||
checkpoint_state_dict.update({'model': model.state_dict()})
|
||||
checkpoint_state_dict.update({'model': model.state_dict(include_optimizer_state)})
|
||||
|
||||
assert os.path.exists(checkpoint_dir), "ERROR: Checkpoint directory doesn't exist: {}".format(checkpoint_dir)
|
||||
|
||||
|
|
@ -625,7 +625,9 @@ class ORTTrainer():
|
|||
if isinstance(model, torch.nn.Module):
|
||||
self.torch_model_ = model
|
||||
self.loss_fn_ = loss_fn
|
||||
self._torch_state_dict_keys = list(model.state_dict().keys())
|
||||
else:
|
||||
self._torch_state_dict_keys = []
|
||||
self.onnx_model_ = model
|
||||
if loss_fn is not None:
|
||||
warnings.warn("loss_fn is not used when creating ORTTrainer because an ONNX model is provided.")
|
||||
|
|
@ -645,8 +647,6 @@ class ORTTrainer():
|
|||
self.world_size = world_size
|
||||
self.use_mixed_precision = use_mixed_precision
|
||||
|
||||
self.original_model_state_keys = list(model.state_dict().keys()) if hasattr(model, 'state_dict') else []
|
||||
|
||||
self.session = None
|
||||
self.device_ = device
|
||||
self.gradient_accumulation_steps = gradient_accumulation_steps
|
||||
|
|
@ -772,7 +772,7 @@ class ORTTrainer():
|
|||
del self.onnx_model_.graph.initializer[w_i]
|
||||
self.onnx_model_.graph.initializer.extend(new_weights)
|
||||
|
||||
def state_dict(self):
|
||||
def state_dict(self, include_optimizer_state=True):
|
||||
if not self.session:
|
||||
warnings.warn("ONNXRuntime training session is not initialized yet. "
|
||||
"Please run train_step or eval_step at least once before calling state_dict().")
|
||||
|
|
@ -790,10 +790,9 @@ class ORTTrainer():
|
|||
torch_state[n.name] = torch.from_numpy(numpy_helper.to_array(n))
|
||||
|
||||
# Need to remove redundant initializers and name suffices to map back to original torch state names
|
||||
torch_state_to_return = {key: torch_state[key] for key in self.original_model_state_keys if key in torch_state} \
|
||||
if self.original_model_state_keys \
|
||||
else torch_state
|
||||
return torch_state_to_return
|
||||
if not include_optimizer_state and self._torch_state_dict_keys:
|
||||
return {key: torch_state[key] for key in self._torch_state_dict_keys if key in torch_state}
|
||||
return torch_state
|
||||
|
||||
def load_state_dict(self, state_dict, strict=False):
|
||||
# Note: It may happen ONNX model has not yet been initialized
|
||||
|
|
|
|||
Loading…
Reference in a new issue