Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion sagemaker-train/src/sagemaker/train/model_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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()

Expand Down
26 changes: 26 additions & 0 deletions sagemaker-train/tests/unit/train/test_model_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
Loading