diff --git a/orttraining/orttraining/python/ort_trainer.py b/orttraining/orttraining/python/ort_trainer.py index 56766ce0d9..0518c23c1d 100644 --- a/orttraining/orttraining/python/ort_trainer.py +++ b/orttraining/orttraining/python/ort_trainer.py @@ -696,6 +696,9 @@ class ORTTrainer(): if self.run_symbolic_shape_infer: self.onnx_model_ = SymbolicShapeInference.infer_shapes(self.onnx_model_, auto_merge=True, guess_output_rank=True) + # old ort session may already exists and occupies GPU memory when creating new session, this may cause OOM error. + # for example, load_state_dict will be called before returing the function, and it calls _init_session again + del self.session self.session, self.train_io_binding, self.eval_io_binding, self.output_name, _, self.output_types = \ create_ort_training_session_with_optimizer( self.onnx_model_, self.device_, diff --git a/orttraining/orttraining/python/training/orttrainer.py b/orttraining/orttraining/python/training/orttrainer.py index 1361a73a90..520bcd1a33 100644 --- a/orttraining/orttraining/python/training/orttrainer.py +++ b/orttraining/orttraining/python/training/orttrainer.py @@ -685,6 +685,9 @@ class ORTTrainer(object): if self.options.utils.run_symbolic_shape_infer: self._onnx_model = SymbolicShapeInference.infer_shapes(self._onnx_model, auto_merge=True, guess_output_rank=True) + # old ort session may already exists and occupies GPU memory when creating new session, this may cause OOM error. + # for example, load_state_dict will be called before returing the function, and it calls _init_session again + del self._training_session # Create training session used by train_step self._create_ort_training_session()