diff --git a/sagemaker-train/src/sagemaker/train/model_trainer.py b/sagemaker-train/src/sagemaker/train/model_trainer.py index 8816ea68f9..7c8c2d32d5 100644 --- a/sagemaker-train/src/sagemaker/train/model_trainer.py +++ b/sagemaker-train/src/sagemaker/train/model_trainer.py @@ -122,7 +122,7 @@ from sagemaker.core.workflow.pipeline_context import PipelineSession, runnable_by_pipeline from sagemaker.core.helper.pipeline_variable import StrPipeVar -from sagemaker.train.local.local_container import _LocalContainer +from sagemaker.train.local.local_container import _LocalContainer, _rmtree class Mode(Enum): @@ -1000,6 +1000,11 @@ def train( environment=training_request["environment"], ) local_container.train(wait) + if self._temp_code_dir is not None: + # The code dir is mounted into the container, which can leave root-owned + # files (e.g. __pycache__) behind that TemporaryDirectory.cleanup() cannot + # remove, so use the same root-aware removal as the container root. + _rmtree(self._temp_code_dir.name, local_container.image, local_container.is_studio) if self._temp_code_dir is not None: self._temp_code_dir.cleanup() diff --git a/sagemaker-train/tests/unit/train/test_model_trainer.py b/sagemaker-train/tests/unit/train/test_model_trainer.py index f990b7ab88..73f771285d 100644 --- a/sagemaker-train/tests/unit/train/test_model_trainer.py +++ b/sagemaker-train/tests/unit/train/test_model_trainer.py @@ -626,6 +626,32 @@ def test_create_input_data_channel_custom_input_s3_key_prefix( assert f"{DEFAULT_BASE_NAME}/input/code" not in mock_upload_data.call_args.kwargs["key_prefix"] +@patch("sagemaker.train.model_trainer._rmtree") +@patch("sagemaker.train.model_trainer._LocalContainer") +def test_local_container_train_removes_temp_code_dir_with_root_owned_files( + mock_local_container, mock_rmtree, tmp_path +): + """The sm_drivers temp dir is mounted into the container, which can leave root-owned + files (e.g. __pycache__) in it, so it must be removed with the root-aware _rmtree.""" + mock_local_container.return_value.image = DEFAULT_IMAGE + mock_local_container.return_value.is_studio = False + + trainer = ModelTrainer( + training_mode=Mode.LOCAL_CONTAINER, + training_image=DEFAULT_IMAGE, + role=DEFAULT_ROLE, + source_code=DEFAULT_SOURCE_CODE, + compute=Compute(instance_type="local", instance_count=1), + local_container_root=str(tmp_path), + ) + + trainer.train(wait=True) + + temp_code_dir = trainer._temp_code_dir.name + mock_local_container.return_value.train.assert_called_once_with(True) + mock_rmtree.assert_called_once_with(temp_code_dir, DEFAULT_IMAGE, False) + + @patch("sagemaker.train.model_trainer.ModelTrainer._resolve_staging_bucket") @patch("sagemaker.train.model_trainer.Session.upload_data") @patch("sagemaker.train.model_trainer.Session.default_bucket")