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
18 changes: 18 additions & 0 deletions migration.md
Original file line number Diff line number Diff line change
Expand Up @@ -231,6 +231,24 @@ train_data = InputData(channel_name="train", data_source="s3://my-bucket/train")
model_trainer.train(input_data_config=[train_data])
```

#### Source code from a Git repository

As in V2 (where `git_config` was passed to the estimator), you can point `SourceCode` at a Git
repository. When `git_config` is set, the repo is cloned at train time and `entry_script`
(and, optionally, `source_dir`) are resolved relative to the clone. `entry_script` is required,
and a local or S3 `source_dir` cannot be combined with `git_config` — when provided, `source_dir`
must be a path relative to the repository root.

```python
source_code = SourceCode(
entry_script="train.py",
git_config={
"repo": "https://github.com/my-org/my-repo.git",
"branch": "main",
},
)
```

### Framework Estimators

**V2 PyTorch:**
Expand Down
7 changes: 7 additions & 0 deletions sagemaker-core/src/sagemaker/core/modules/configs.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,12 @@ class SourceCode(BaseConfig):
Only applicable when ``command`` is provided; top-level ``hyperparameters`` are not
passed as CLI arguments in ``command`` mode -- they are available inside the container
via the ``SM_HPS`` environment variable.
git_config (Optional[dict]):
Git configuration used to clone the repository that contains the source code, including
``repo``, ``branch``, ``commit``, ``2FA_enabled``, ``username``, ``password`` and
``token``. Only ``repo`` is required. When provided, the repository is cloned and
``source_dir``/``entry_script`` are resolved relative to the clone; a local
``source_dir`` cannot be combined with ``git_config``.
ignore_patterns: (Optional[List[str]]) :
The ignore patterns to ignore specific files/folders when uploading to S3. If not specified,
default to: ['.env', '.git', '__pycache__', '.DS_Store', '.cache', '.ipynb_checkpoints'].
Expand All @@ -121,6 +127,7 @@ class SourceCode(BaseConfig):
entry_script: Optional[str] = None
command: Optional[str] = None
args: Optional[List[Union[str, int, float]]] = None
git_config: Optional[dict] = None
ignore_patterns: Optional[List[str]] = [
".env",
".git",
Expand Down
7 changes: 7 additions & 0 deletions sagemaker-core/src/sagemaker/core/training/configs.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,6 +114,12 @@ class SourceCode(BaseConfig):
Only applicable when ``command`` is provided; top-level ``hyperparameters`` are not
passed as CLI arguments in ``command`` mode -- they are available inside the container
via the ``SM_HPS`` environment variable.
git_config (Optional[dict]):
Git configuration used to clone the repository that contains the source code, including
``repo``, ``branch``, ``commit``, ``2FA_enabled``, ``username``, ``password`` and
``token``. Only ``repo`` is required. When provided, the repository is cloned and
``source_dir``/``entry_script`` are resolved relative to the clone; a local
``source_dir`` cannot be combined with ``git_config``.
ignore_patterns: (Optional[List[str]]) :
The ignore patterns to ignore specific files/folders when uploading to S3. If not specified,
default to: ['.env', '.git', '__pycache__', '.DS_Store', '.cache', '.ipynb_checkpoints'].
Expand All @@ -124,6 +130,7 @@ class SourceCode(BaseConfig):
entry_script: Optional[StrPipeVar] = None
command: Optional[StrPipeVar] = None
args: Optional[List[Union[str, int, float]]] = None
git_config: Optional[dict] = None
ignore_patterns: Optional[List[str]] = [
".env",
".git",
Expand Down
61 changes: 56 additions & 5 deletions sagemaker-train/src/sagemaker/train/model_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,7 @@
_is_valid_s3_uri,
safe_serialize,
)
from sagemaker.core.git_utils import git_clone_repo, _validate_git_config
from sagemaker.train.types import DataSourceType
from sagemaker.train.constants import (
SM_CODE,
Expand Down Expand Up @@ -540,6 +541,31 @@ def _validate_distributed_config(
def _validate_source_code(self, source_code: Optional[SourceCode]):
"""Validate the source code configuration."""
if source_code:
if source_code.git_config:
# When ``git_config`` is provided, ``source_dir``/``entry_script``/``requirements``
# are resolved relative to the cloned repository at train time.
if not source_code.entry_script:
raise ValueError(
"'entry_script' must be provided in 'source_code' when 'git_config' is "
"set. It is resolved relative to the cloned Git repository."
)
# A local or S3 ``source_dir`` is mutually exclusive with ``git_config``: with
# ``git_config`` the source lives in the repo, so ``source_dir`` (when given) must
# be a path relative to the repository root -- not an absolute local path, an S3
# URI, or a packaged tar.gz.
source_dir = source_code.source_dir
if source_dir and (
os.path.isabs(source_dir)
or source_dir.lower().startswith("s3://")
or source_dir.endswith(".tar.gz")
):
raise ValueError(
"'git_config' and a local or S3 'source_dir' are mutually exclusive. "
"When 'git_config' is provided, 'source_dir' must be a relative path "
"to a directory within the Git repository."
)
_validate_git_config(source_code.git_config)
return
if source_code.requirements or source_code.entry_script:
source_dir = source_code.source_dir
requirements = source_code.requirements
Expand Down Expand Up @@ -772,29 +798,54 @@ def _create_training_job_args(
driver_dir = os.path.join(self._temp_code_dir.name, "distributed_drivers")
shutil.copytree(distributed_driver_dir, driver_dir, dirs_exist_ok=True)

# Work on a copy so a git_config clone does not permanently mutate the user's
# ``source_code`` (``train()`` is re-callable and must re-clone on each call).
source_code = self.source_code

# If git_config is provided, clone the repository and resolve source_dir/entry_script
# to point at the local clone so the source code channel and train script use it.
if source_code.git_config:
source_code = source_code.model_copy(deep=True)
updated_paths = git_clone_repo(
git_config=source_code.git_config,
entry_point=source_code.entry_script,
source_dir=source_code.source_dir,
dependencies=None,
)
if updated_paths["source_dir"]:
source_code.source_dir = updated_paths["source_dir"]
else:
# No source_dir was provided: entry_point was resolved to an absolute path
# within the clone. Use its parent directory as the source directory.
source_code.source_dir = os.path.dirname(updated_paths["entry_point"])
source_code.entry_script = os.path.basename(updated_paths["entry_point"])
# Drop git_config now that the repo is cloned so credentials are never serialized
# into the source_code.json that is uploaded to the container.
source_code.git_config = None

# If source code is provided, create a channel for the source code
# The source code will be mounted at /opt/ml/input/data/code in the container
if self.source_code.source_dir:
if source_code.source_dir:
source_code_channel = self.create_input_data_channel(
channel_name=SM_CODE,
data_source=self.source_code.source_dir,
data_source=source_code.source_dir,
key_prefix=input_data_key_prefix,
ignore_patterns=self.source_code.ignore_patterns,
ignore_patterns=source_code.ignore_patterns,
instance_group_names=managed_channel_instance_group_names,
)
final_input_data_config.append(source_code_channel)

self._prepare_train_script(
tmp_dir=self._temp_code_dir,
source_code=self.source_code,
source_code=source_code,
distributed=self.distributed,
)

if isinstance(self.distributed, Torchrun) and self.distributed.smp:
mp_parameters = self.distributed.smp._to_mp_hyperparameters()
string_hyper_parameters.update(mp_parameters)

self._write_source_code_json(tmp_dir=self._temp_code_dir, source_code=self.source_code)
self._write_source_code_json(tmp_dir=self._temp_code_dir, source_code=source_code)
self._write_distributed_json(tmp_dir=self._temp_code_dir, distributed=self.distributed)

# Create an input channel for drivers packaged by the sdk.
Expand Down
217 changes: 217 additions & 0 deletions sagemaker-train/tests/unit/train/test_model_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -2369,3 +2369,220 @@ def test_output_data_strips_trailing_slash_on_fallback(model_trainer):
assert model_trainer.output_data == (
f"s3://{DEFAULT_BUCKET}/{DEFAULT_BUCKET_PREFIX}" "/my-training-job/output/output.tar.gz"
)


# ---------------------------------------------------------------------------
# git_config support in SourceCode / ModelTrainer (issue #5571)
# ---------------------------------------------------------------------------

GIT_CONFIG = {"repo": "https://github.com/example/repo.git", "branch": "main"}


def test_source_code_accepts_git_config_train_config():
"""SourceCode (sagemaker.train.configs) exposes an optional ``git_config`` field."""
source_code = SourceCode(entry_script="train.py", git_config=GIT_CONFIG)
assert source_code.git_config == GIT_CONFIG


def test_source_code_accepts_git_config_core_modules_config():
"""SourceCode (sagemaker.core.modules.configs) exposes an optional ``git_config`` field."""
from sagemaker.core.modules.configs import SourceCode as CoreSourceCode

source_code = CoreSourceCode(entry_script="train.py", git_config=GIT_CONFIG)
assert source_code.git_config == GIT_CONFIG


@patch("sagemaker.train.model_trainer.TrainingJob")
@patch("sagemaker.train.model_trainer.ModelTrainer.create_input_data_channel")
@patch("sagemaker.train.model_trainer.git_clone_repo")
def test_source_code_git_config_clones_repo(
mock_git_clone_repo, mock_create_input_data_channel, mock_training_job, modules_session
):
"""When ``git_config`` is set, the repo is cloned and ``source_dir`` points at the clone."""
clone_dir = tempfile.mkdtemp()
entry = "custom_script.py"
with open(os.path.join(clone_dir, entry), "w") as f:
f.write("print('hello')\n")

# Mirror git_clone_repo's contract when no source_dir is provided: entry_point is
# resolved to an absolute path inside the clone and source_dir stays None.
mock_git_clone_repo.return_value = {
"entry_point": os.path.join(clone_dir, entry),
"source_dir": None,
"dependencies": None,
}

trainer = ModelTrainer(
training_image=DEFAULT_IMAGE,
role=DEFAULT_ROLE,
compute=DEFAULT_COMPUTE_CONFIG,
stopping_condition=DEFAULT_STOPPING_CONDITION,
output_data_config=DEFAULT_OUTPUT_DATA_CONFIG,
source_code=SourceCode(entry_script=entry, git_config=GIT_CONFIG),
sagemaker_session=modules_session,
)

trainer.train()

mock_git_clone_repo.assert_called_once()
call = mock_git_clone_repo.call_args
forwarded_git_config = call.kwargs.get("git_config", call.args[0] if call.args else None)
assert forwarded_git_config == GIT_CONFIG

# The source_code channel points at the resolved local clone so entry_script/requirements
# are found there.
calls_by_channel = {
c.kwargs.get("channel_name"): c.kwargs
for c in mock_create_input_data_channel.call_args_list
}
assert SM_CODE in calls_by_channel, "source_dir channel was not created from the clone"
assert calls_by_channel[SM_CODE].get("data_source") == clone_dir

# The user's SourceCode must NOT be mutated: train() is re-callable and must re-clone.
assert trainer.source_code.source_dir is None
assert trainer.source_code.git_config == GIT_CONFIG


@patch("sagemaker.train.model_trainer.TrainingJob")
@patch("sagemaker.train.model_trainer.ModelTrainer.create_input_data_channel")
@patch("sagemaker.train.model_trainer.git_clone_repo")
def test_source_code_git_config_with_relative_source_dir(
mock_git_clone_repo, mock_create_input_data_channel, mock_training_job, modules_session
):
"""A relative ``source_dir`` (a path within the repo) is forwarded and then overridden."""
clone_source_dir = tempfile.mkdtemp()
mock_git_clone_repo.return_value = {
"entry_point": "train.py",
"source_dir": clone_source_dir,
"dependencies": None,
}

trainer = ModelTrainer(
training_image=DEFAULT_IMAGE,
role=DEFAULT_ROLE,
compute=DEFAULT_COMPUTE_CONFIG,
stopping_condition=DEFAULT_STOPPING_CONDITION,
output_data_config=DEFAULT_OUTPUT_DATA_CONFIG,
source_code=SourceCode(source_dir="src", entry_script="train.py", git_config=GIT_CONFIG),
sagemaker_session=modules_session,
)

trainer.train()

mock_git_clone_repo.assert_called_once()
call = mock_git_clone_repo.call_args
forwarded_source_dir = call.kwargs.get("source_dir", None)
assert forwarded_source_dir == "src"

# After cloning, the source_code channel uses the resolved clone location.
calls_by_channel = {
c.kwargs.get("channel_name"): c.kwargs
for c in mock_create_input_data_channel.call_args_list
}
assert calls_by_channel[SM_CODE].get("data_source") == clone_source_dir

# The user's SourceCode must NOT be mutated.
assert trainer.source_code.source_dir == "src"
assert trainer.source_code.git_config == GIT_CONFIG


@patch("sagemaker.train.model_trainer.TrainingJob")
@patch("sagemaker.train.model_trainer.ModelTrainer.create_input_data_channel")
@patch("sagemaker.train.model_trainer.git_clone_repo")
def test_git_config_credentials_not_serialized(
mock_git_clone_repo, mock_create_input_data_channel, mock_training_job, modules_session
):
"""git_config (and any embedded credentials) must not be written into source_code.json."""
clone_dir = tempfile.mkdtemp()
entry = "train.py"
with open(os.path.join(clone_dir, entry), "w") as f:
f.write("x = 1\n")
mock_git_clone_repo.return_value = {
"entry_point": os.path.join(clone_dir, entry),
"source_dir": None,
"dependencies": None,
}
git_config = {"repo": "https://github.com/example/repo.git", "token": "super-secret-token"}

trainer = ModelTrainer(
training_image=DEFAULT_IMAGE,
role=DEFAULT_ROLE,
compute=DEFAULT_COMPUTE_CONFIG,
stopping_condition=DEFAULT_STOPPING_CONDITION,
output_data_config=DEFAULT_OUTPUT_DATA_CONFIG,
source_code=SourceCode(entry_script=entry, git_config=git_config),
sagemaker_session=modules_session,
)

with patch.object(ModelTrainer, "_write_source_code_json") as mock_write:
trainer.train()

write_call = mock_write.call_args
written_source_code = write_call.kwargs.get(
"source_code", write_call.args[1] if len(write_call.args) > 1 else None
)
assert written_source_code.git_config is None


def test_git_config_and_local_source_dir_are_mutually_exclusive(modules_session):
"""A local (absolute) ``source_dir`` cannot be combined with ``git_config``."""
with pytest.raises(ValueError, match="mutually exclusive"):
ModelTrainer(
training_image=DEFAULT_IMAGE,
role=DEFAULT_ROLE,
compute=DEFAULT_COMPUTE_CONFIG,
stopping_condition=DEFAULT_STOPPING_CONDITION,
output_data_config=DEFAULT_OUTPUT_DATA_CONFIG,
source_code=SourceCode(
source_dir=os.path.abspath(DEFAULT_SOURCE_DIR),
entry_script="custom_script.py",
git_config=GIT_CONFIG,
),
sagemaker_session=modules_session,
)


def test_git_config_and_s3_source_dir_are_mutually_exclusive(modules_session):
"""An S3 ``source_dir`` cannot be combined with ``git_config``."""
with pytest.raises(ValueError, match="mutually exclusive"):
ModelTrainer(
training_image=DEFAULT_IMAGE,
role=DEFAULT_ROLE,
compute=DEFAULT_COMPUTE_CONFIG,
stopping_condition=DEFAULT_STOPPING_CONDITION,
output_data_config=DEFAULT_OUTPUT_DATA_CONFIG,
source_code=SourceCode(
source_dir="s3://bucket/code/",
entry_script="custom_script.py",
git_config=GIT_CONFIG,
),
sagemaker_session=modules_session,
)


def test_git_config_requires_entry_script(modules_session):
"""``entry_script`` is required when ``git_config`` is provided (command-only is rejected)."""
with pytest.raises(ValueError, match="entry_script"):
ModelTrainer(
training_image=DEFAULT_IMAGE,
role=DEFAULT_ROLE,
compute=DEFAULT_COMPUTE_CONFIG,
stopping_condition=DEFAULT_STOPPING_CONDITION,
output_data_config=DEFAULT_OUTPUT_DATA_CONFIG,
source_code=SourceCode(command="python train.py", git_config=GIT_CONFIG),
sagemaker_session=modules_session,
)


def test_git_config_without_repo_raises(modules_session):
"""``git_config`` must contain a ``repo`` key."""
with pytest.raises(ValueError, match="repo"):
ModelTrainer(
training_image=DEFAULT_IMAGE,
role=DEFAULT_ROLE,
compute=DEFAULT_COMPUTE_CONFIG,
stopping_condition=DEFAULT_STOPPING_CONDITION,
output_data_config=DEFAULT_OUTPUT_DATA_CONFIG,
source_code=SourceCode(entry_script="train.py", git_config={"branch": "main"}),
sagemaker_session=modules_session,
)
Loading