Uh oh!
There was an error while loading. Please reload this page.
- Notifications
You must be signed in to change notification settings - Fork 1.3k
fix: Add additional dependencies for ModelTrainer (5668)#5731
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base:master
Are you sure you want to change the base?
Uh oh!
There was an error while loading. Please reload this page.
Changes from all commits
File filter
Filter by extension
Conversations
Uh oh!
There was an error while loading. Please reload this page.
Jump to
Uh oh!
There was an error while loading. Please reload this page.
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -63,6 +63,9 @@ | ||
| "amazon.nova-pro-v1:0": ["us-east-1"] | ||
| } | ||
aviruthen marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. | ||
| SM_DEPENDENCIES = "sm_dependencies" | ||
| SM_DEPENDENCIES_CONTAINER_PATH = "/opt/ml/input/data/sm_dependencies" | ||
| SM_RECIPE = "recipe" | ||
| SM_RECIPE_YAML = "recipe.yaml" | ||
| SM_RECIPE_CONTAINER_PATH = f"/opt/ml/input/data/recipe/{SM_RECIPE_YAML}" | ||
| SM_RECIPE_CONTAINER_PATH = f"/opt/ml/input/data/recipe/{SM_RECIPE_YAML}" | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -83,6 +83,8 @@ | ||
| SM_CODE_CONTAINER_PATH, | ||
| SM_DRIVERS, | ||
| SM_DRIVERS_LOCAL_PATH, | ||
| SM_DEPENDENCIES, | ||
| SM_DEPENDENCIES_CONTAINER_PATH, | ||
| SM_RECIPE, | ||
| SM_RECIPE_YAML, | ||
| SM_RECIPE_CONTAINER_PATH, | ||
| @@ -99,6 +101,7 @@ | ||
| EXECUTE_BASIC_SCRIPT_DRIVER, | ||
| INSTALL_AUTO_REQUIREMENTS, | ||
| INSTALL_REQUIREMENTS, | ||
| INSTALL_DEPENDENCIES, | ||
| ) | ||
| from sagemaker.core.telemetry.telemetry_logging import _telemetry_emitter | ||
| from sagemaker.core.telemetry.constants import Feature | ||
| @@ -269,6 +272,7 @@ class ModelTrainer(BaseModel): | ||
| # Private Attributes for AWS_Batch | ||
| _temp_code_dir: Optional[TemporaryDirectory] = PrivateAttr(default=None) | ||
| _temp_deps_dir: Optional[TemporaryDirectory] = PrivateAttr(default=None) | ||
| CONFIGURABLE_ATTRIBUTES: ClassVar[List[str]] = [ | ||
| "role", | ||
| @@ -408,6 +412,8 @@ def __del__(self): | ||
| self._temp_recipe_train_dir.cleanup() | ||
| if self._temp_code_dir is not None: | ||
| self._temp_code_dir.cleanup() | ||
| if self._temp_deps_dir is not None: | ||
| self._temp_deps_dir.cleanup() | ||
| def _validate_training_image_and_algorithm_name( | ||
| self, training_image: Optional[str], algorithm_name: Optional[str] | ||
| @@ -484,6 +490,13 @@ def _validate_source_code(self, source_code: Optional[SourceCode]): | ||
| f"Invalid 'entry_script': {entry_script}. " | ||
| "Must be a valid file within the 'source_dir'.", | ||
| ) | ||
aviruthen marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. | ||
| if source_code.dependencies: | ||
| for dep_path in source_code.dependencies: | ||
| if not _is_valid_path(dep_path, path_type="Directory"): | ||
| raise ValueError( | ||
| f"Invalid dependency path: {dep_path}. " | ||
| "Each dependency must be a valid local directory path." | ||
| ) | ||
| @staticmethod | ||
| def _validate_and_fetch_hyperparameters_file(hyperparameters_file: str): | ||
| @@ -654,6 +667,21 @@ def _create_training_job_args( | ||
| ) | ||
| final_input_data_config.append(source_code_channel) | ||
| # If dependencies are provided, create a channel for the dependencies | ||
| # The dependencies will be mounted at /opt/ml/input/data/sm_dependencies | ||
| if self.source_code.dependencies: | ||
| self._temp_deps_dir = TemporaryDirectory() | ||
| for dep_path in self.source_code.dependencies: | ||
| dep_basename = os.path.basename(os.path.normpath(dep_path)) | ||
| dest_path = os.path.join(self._temp_deps_dir.name, dep_basename) | ||
| shutil.copytree(dep_path, dest_path, dirs_exist_ok=True) | ||
| dependencies_channel = self.create_input_data_channel( | ||
| channel_name=SM_DEPENDENCIES, | ||
| data_source=self._temp_deps_dir.name, | ||
| key_prefix=input_data_key_prefix, | ||
| ) | ||
| final_input_data_config.append(dependencies_channel) | ||
| self._prepare_train_script( | ||
| tmp_dir=self._temp_code_dir, | ||
| source_code=self.source_code, | ||
| @@ -813,6 +841,9 @@ def train( | ||
| local_container.train(wait) | ||
| if self._temp_code_dir is not None: | ||
| self._temp_code_dir.cleanup() | ||
| if self._temp_deps_dir is not None: | ||
| self._temp_deps_dir.cleanup() | ||
| self._temp_deps_dir = None | ||
| def create_input_data_channel( | ||
| @@ -1010,6 +1041,10 @@ def _prepare_train_script( | ||
| base_command = source_code.command.split() | ||
| base_command = " ".join(base_command) | ||
| install_dependencies = "" | ||
aviruthen marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. | ||
| if source_code.dependencies: | ||
| install_dependencies = INSTALL_DEPENDENCIES | ||
| install_requirements = "" | ||
| if source_code.requirements: | ||
| if self._jumpstart_config and source_code.requirements == "auto": | ||
| @@ -1049,6 +1084,7 @@ def _prepare_train_script( | ||
| train_script = TRAIN_SCRIPT_TEMPLATE.format( | ||
| working_dir=working_dir, | ||
| install_dependencies=install_dependencies, | ||
| install_requirements=install_requirements, | ||
| execute_driver=execute_driver, | ||
| ) | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -39,6 +39,29 @@ | ||
| $SM_PIP_CMD install -r {requirements_file} | ||
| """ | ||
| INSTALL_DEPENDENCIES = """ | ||
aviruthen marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. | ||
| echo "Setting up additional dependencies" | ||
| if [ -d /opt/ml/input/data/sm_dependencies ]; then | ||
| for dep in /opt/ml/input/data/sm_dependencies/*; do | ||
| if [ -d "$dep" ]; then | ||
| echo "Adding directory $dep to PYTHONPATH" | ||
| export PYTHONPATH="$dep:$PYTHONPATH" | ||
| elif [ -f "$dep" ]; then | ||
| case "$dep" in | ||
| *.whl|*.tar.gz) | ||
| echo "Installing package $dep via pip" | ||
| $SM_PIP_CMD install "$dep" | ||
| ;; | ||
| *) | ||
| echo "Adding parent directory of $dep to PYTHONPATH" | ||
| export PYTHONPATH="/opt/ml/input/data/sm_dependencies:$PYTHONPATH" | ||
| ;; | ||
| esac | ||
| fi | ||
| done | ||
| fi | ||
| """ | ||
| EXEUCTE_DISTRIBUTED_DRIVER = """ | ||
| echo "Running {driver_name} Driver" | ||
| $SM_PYTHON_CMD /opt/ml/input/data/sm_drivers/distributed_drivers/{driver_script} | ||
| @@ -95,6 +118,7 @@ | ||
| set -x | ||
| {working_dir} | ||
| {install_dependencies} | ||
| {install_requirements} | ||
| {execute_driver} | ||
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.