From 19269d94f87fb505fdf96f87369df6b3bbc30a69 Mon Sep 17 00:00:00 2001 From: Mohamed Zeidan Date: Mon, 28 Sep 2026 14:27:38 -0700 Subject: [PATCH] feat: support git_config in SourceCode / ModelTrainer (#5571) --- migration.md | 18 ++ .../src/sagemaker/core/modules/configs.py | 7 + .../src/sagemaker/core/training/configs.py | 7 + .../src/sagemaker/train/model_trainer.py | 61 ++++- .../tests/unit/train/test_model_trainer.py | 217 ++++++++++++++++++ 5 files changed, 305 insertions(+), 5 deletions(-) diff --git a/migration.md b/migration.md index 0ea3e9bef1..a3e3c0184f 100644 --- a/migration.md +++ b/migration.md @@ -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:** diff --git a/sagemaker-core/src/sagemaker/core/modules/configs.py b/sagemaker-core/src/sagemaker/core/modules/configs.py index 865018f50c..d2e9a51641 100644 --- a/sagemaker-core/src/sagemaker/core/modules/configs.py +++ b/sagemaker-core/src/sagemaker/core/modules/configs.py @@ -104,6 +104,12 @@ class SourceCode(BaseConfig): command (Optional[str]): The command(s) to execute in the training job container. Example: "python my_script.py". If not specified, entry_script must be provided. + 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']. @@ -113,6 +119,7 @@ class SourceCode(BaseConfig): requirements: Optional[str] = None entry_script: Optional[str] = None command: Optional[str] = None + git_config: Optional[dict] = None ignore_patterns: Optional[List[str]] = [ ".env", ".git", diff --git a/sagemaker-core/src/sagemaker/core/training/configs.py b/sagemaker-core/src/sagemaker/core/training/configs.py index 17e40b0771..208611af77 100644 --- a/sagemaker-core/src/sagemaker/core/training/configs.py +++ b/sagemaker-core/src/sagemaker/core/training/configs.py @@ -107,6 +107,12 @@ class SourceCode(BaseConfig): command (Optional[StrPipeVar]): The command(s) to execute in the training job container. Example: "python my_script.py". If not specified, entry_script must be provided. + 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']. @@ -116,6 +122,7 @@ class SourceCode(BaseConfig): requirements: Optional[StrPipeVar] = None entry_script: Optional[StrPipeVar] = None command: Optional[StrPipeVar] = None + git_config: Optional[dict] = None ignore_patterns: Optional[List[str]] = [ ".env", ".git", diff --git a/sagemaker-train/src/sagemaker/train/model_trainer.py b/sagemaker-train/src/sagemaker/train/model_trainer.py index 35b1395cce..510b3f075b 100644 --- a/sagemaker-train/src/sagemaker/train/model_trainer.py +++ b/sagemaker-train/src/sagemaker/train/model_trainer.py @@ -80,6 +80,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, @@ -480,6 +481,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 @@ -709,21 +735,46 @@ 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, ) @@ -731,7 +782,7 @@ def _create_training_job_args( 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. diff --git a/sagemaker-train/tests/unit/train/test_model_trainer.py b/sagemaker-train/tests/unit/train/test_model_trainer.py index ceb371dfa6..571cf3a95b 100644 --- a/sagemaker-train/tests/unit/train/test_model_trainer.py +++ b/sagemaker-train/tests/unit/train/test_model_trainer.py @@ -2220,3 +2220,220 @@ def test_log_actionable_client_error_other_codes_stay_silent(caplog): _log_actionable_client_error(error) assert caplog.text == "" + + +# --------------------------------------------------------------------------- +# 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, + )