Buckets:
hf-doc-build/doc-dev / sagemaker /pr_2709 /en /tutorials /sagemaker-sdk /training-sagemaker-sdk.html
| <meta charset="utf-8" /><meta name="hf:doc:metadata" content="{"title":"Train models on Amazon SageMaker with the SageMaker SDK","local":"train-models-on-amazon-sagemaker-with-the-sagemaker-sdk","sections":[{"title":"Prepare a training script","local":"prepare-a-training-script","sections":[],"depth":2},{"title":"Create a ModelTrainer","local":"create-a-modeltrainer","sections":[],"depth":2},{"title":"Start the training job","local":"start-the-training-job","sections":[],"depth":2},{"title":"Training output and checkpoints","local":"training-output-and-checkpoints","sections":[],"depth":2},{"title":"Access the trained model","local":"access-the-trained-model","sections":[],"depth":2},{"title":"Distributed training","local":"distributed-training","sections":[{"title":"Data parallelism","local":"data-parallelism","sections":[],"depth":3},{"title":"Model parallelism","local":"model-parallelism","sections":[],"depth":3}],"depth":2},{"title":"Spot instances","local":"spot-instances","sections":[],"depth":2},{"title":"Git repository","local":"git-repository","sections":[],"depth":2},{"title":"SageMaker metrics","local":"sagemaker-metrics","sections":[],"depth":2},{"title":"What’s next","local":"whats-next","sections":[],"depth":2}],"depth":1}"/> | |
| <link href="/docs/sagemaker/pr_2709/en/_app/immutable/entry/start.COiK-bUV.js" rel="modulepreload"> | |
| <link href="/docs/sagemaker/pr_2709/en/_app/immutable/chunks/UGDyBJsP.js" rel="modulepreload"> | |
| <link href="/docs/sagemaker/pr_2709/en/_app/immutable/chunks/Paddn1qZ.js" rel="modulepreload"> | |
| <link href="/docs/sagemaker/pr_2709/en/_app/immutable/entry/app.CjYANwwM.js" rel="modulepreload"> | |
| <link href="/docs/sagemaker/pr_2709/en/_app/immutable/chunks/Beft82ym.js" rel="modulepreload"> | |
| <link href="/docs/sagemaker/pr_2709/en/_app/immutable/chunks/Ds9Xs-bq.js" rel="modulepreload"> | |
| <link href="/docs/sagemaker/pr_2709/en/_app/immutable/chunks/DsnmJJEf.js" rel="modulepreload"> | |
| <link href="/docs/sagemaker/pr_2709/en/_app/immutable/chunks/DnXYL5Yp.js" rel="modulepreload"> | |
| <link href="/docs/sagemaker/pr_2709/en/_app/immutable/nodes/0.CsY3iCn7.js" rel="modulepreload"> | |
| <link href="/docs/sagemaker/pr_2709/en/_app/immutable/chunks/CxJVP3KE.js" rel="modulepreload"> | |
| <link href="/docs/sagemaker/pr_2709/en/_app/immutable/nodes/30.BO7b6Mz0.js" rel="modulepreload"> | |
| <link href="/docs/sagemaker/pr_2709/en/_app/immutable/chunks/Cz9gzD1e.js" rel="modulepreload"> | |
| <link href="/docs/sagemaker/pr_2709/en/_app/immutable/chunks/CFAwabh-.js" rel="modulepreload"> | |
| <!--f7qhpl--><meta name="hf:doc:metadata" content="{"title":"Train models on Amazon SageMaker with the SageMaker SDK","local":"train-models-on-amazon-sagemaker-with-the-sagemaker-sdk","sections":[{"title":"Prepare a training script","local":"prepare-a-training-script","sections":[],"depth":2},{"title":"Create a ModelTrainer","local":"create-a-modeltrainer","sections":[],"depth":2},{"title":"Start the training job","local":"start-the-training-job","sections":[],"depth":2},{"title":"Training output and checkpoints","local":"training-output-and-checkpoints","sections":[],"depth":2},{"title":"Access the trained model","local":"access-the-trained-model","sections":[],"depth":2},{"title":"Distributed training","local":"distributed-training","sections":[{"title":"Data parallelism","local":"data-parallelism","sections":[],"depth":3},{"title":"Model parallelism","local":"model-parallelism","sections":[],"depth":3}],"depth":2},{"title":"Spot instances","local":"spot-instances","sections":[],"depth":2},{"title":"Git repository","local":"git-repository","sections":[],"depth":2},{"title":"SageMaker metrics","local":"sagemaker-metrics","sections":[],"depth":2},{"title":"What’s next","local":"whats-next","sections":[],"depth":2}],"depth":1}"/><!----> | |
| <link href="/docs/sagemaker/pr_2709/en/_app/immutable/assets/0.tn0RQdqM.css" rel="modulepreload"> <!--[--><!--[0--><!--[--><!--[0--><!--[--><p></p> <div class="items-center shrink-0 min-w-[100px] max-sm:min-w-[50px] justify-end ml-auto flex" style="float: right; margin-left: 10px; display: inline-flex; position: relative; z-index: 10;"><div class="inline-flex rounded-md max-sm:rounded-sm"><button class="inline-flex items-center gap-1 h-7 max-sm:h-7 px-2 max-sm:px-1.5 text-sm font-medium text-gray-800 border border-r-0 rounded-l-md max-sm:rounded-l-sm border-gray-200 bg-white hover:shadow-inner dark:border-gray-850 dark:bg-gray-950 dark:text-gray-200 dark:hover:bg-gray-800" aria-live="polite"><span class="inline-flex items-center justify-center rounded-md p-0.5 max-sm:p-0 hover:text-gray-800 dark:hover:text-gray-200"><svg class="sm:size-3.5 size-3" xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----></span> <span>Copy page</span></button> <button class="inline-flex items-center justify-center w-6 max-sm:w-5 h-7 max-sm:h-7 disabled:pointer-events-none text-sm text-gray-500 hover:text-gray-700 dark:hover:text-white rounded-r-md max-sm:rounded-r-sm border border-l transition border-gray-200 bg-white hover:shadow-inner dark:border-gray-850 dark:bg-gray-950 dark:text-gray-200 dark:hover:bg-gray-800" aria-haspopup="menu" aria-expanded="false" aria-label="Open copy menu"><svg class="transition-transform text-gray-400 overflow-visible sm:size-3.5 size-3 rotate-0" width="1em" height="1em" viewBox="0 0 12 7" fill="none" xmlns="http://www.w3.org/2000/svg"><path d="M1 1L6 6L11 1" stroke="currentColor"></path></svg><!----></button></div> <!--[-1--><!--]--></div><!----> <!--[0--><h1 class="relative group"><a id="train-models-on-amazon-sagemaker-with-the-sagemaker-sdk" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#train-models-on-amazon-sagemaker-with-the-sagemaker-sdk"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Train models on Amazon SageMaker with the SageMaker SDK</span></h1><!--]--><!----> <p>This guide shows how to train models with the SageMaker Python SDK <code>ModelTrainer</code> and your own training script. Make sure you have <a href="./setup-sagemaker-sdk">set up the SageMaker SDK</a> first.</p> <p>The examples come from the <a href="../../examples/sagemaker-sdk-fine-tune-llm-sft">Fine-Tune an LLM with the SageMaker SDK and TRL</a> example, which fine-tunes <code>Qwen/Qwen3-0.6B</code> with TRL <code>SFTTrainer</code> on the Hugging Face PyTorch training DLC. The example is the complete runnable version; this guide explains each concept.</p> <div class="mermaid-chart " style="text-align: center;"></div><!----> <p>Learn how to:</p> <ul><li><a href="#prepare-a-training-script">Prepare a training script</a>.</li> <li><a href="#create-a-modeltrainer">Create a ModelTrainer</a>.</li> <li><a href="#start-the-training-job">Start the training job</a>.</li> <li><a href="#training-output-and-checkpoints">Manage training output and checkpoints</a>.</li> <li><a href="#access-the-trained-model">Access the trained model</a>.</li> <li><a href="#distributed-training">Scale to distributed training</a>.</li> <li><a href="#spot-instances">Save with spot instances</a>.</li> <li><a href="#git-repository">Load a training script from a GitHub repository</a>.</li> <li><a href="#sagemaker-metrics">Collect training metrics</a>.</li></ul> <!--[1--><h2 class="relative group"><a id="prepare-a-training-script" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#prepare-a-training-script"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Prepare a training script</span></h2><!--]--><!----> <p>A SageMaker training script is a regular Python script that reads two things from its environment: hyperparameters as command-line arguments, and directory locations as environment variables. The most useful environment variables (see the <a href="https://github.com/aws/sagemaker-training-toolkit/blob/master/ENVIRONMENT_VARIABLES.md" rel="nofollow">full list</a>):</p> <ul><li><code>SM_MODEL_DIR</code>: the directory the job uploads to S3 as <code>model.tar.gz</code> when training finishes. Always <code>/opt/ml/model</code>.</li> <li><code>SM_NUM_GPUS</code>: the number of GPUs available on the instance.</li> <li><code>SM_CHANNEL_XXXX</code>: the path to the input data for channel <code>XXXX</code> when you pass data channels (see <a href="#start-the-training-job">Start the training job</a>).</li></ul> <p>The notebook’s <code>scripts/train.py</code>:</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-python "><!----><span class="hljs-keyword">import</span> argparse | |
| <span class="hljs-keyword">import</span> os | |
| <span class="hljs-keyword">from</span> datasets <span class="hljs-keyword">import</span> load_dataset | |
| <span class="hljs-keyword">from</span> trl <span class="hljs-keyword">import</span> SFTConfig, SFTTrainer | |
| <span class="hljs-keyword">def</span> <span class="hljs-title function_">parse_args</span>(): | |
| parser = argparse.ArgumentParser() | |
| <span class="hljs-comment"># hyperparameters sent by the ModelTrainer arrive as command-line arguments</span> | |
| parser.add_argument(<span class="hljs-string">"--model_name"</span>, <span class="hljs-built_in">type</span>=<span class="hljs-built_in">str</span>, default=<span class="hljs-string">"Qwen/Qwen3-0.6B"</span>) | |
| parser.add_argument(<span class="hljs-string">"--dataset_name"</span>, <span class="hljs-built_in">type</span>=<span class="hljs-built_in">str</span>, default=<span class="hljs-string">"trl-lib/Capybara"</span>) | |
| parser.add_argument(<span class="hljs-string">"--max_steps"</span>, <span class="hljs-built_in">type</span>=<span class="hljs-built_in">int</span>, default=<span class="hljs-number">50</span>) | |
| parser.add_argument(<span class="hljs-string">"--train_batch_size"</span>, <span class="hljs-built_in">type</span>=<span class="hljs-built_in">int</span>, default=<span class="hljs-number">4</span>) | |
| parser.add_argument(<span class="hljs-string">"--learning_rate"</span>, <span class="hljs-built_in">type</span>=<span class="hljs-built_in">float</span>, default=<span class="hljs-number">2e-5</span>) | |
| <span class="hljs-comment"># SageMaker directories: SM_MODEL_DIR is archived to S3 as model.tar.gz</span> | |
| parser.add_argument(<span class="hljs-string">"--model_dir"</span>, <span class="hljs-built_in">type</span>=<span class="hljs-built_in">str</span>, default=os.environ[<span class="hljs-string">"SM_MODEL_DIR"</span>]) | |
| parser.add_argument(<span class="hljs-string">"--output_dir"</span>, <span class="hljs-built_in">type</span>=<span class="hljs-built_in">str</span>, default=os.environ.get(<span class="hljs-string">"SM_OUTPUT_DATA_DIR"</span>, <span class="hljs-string">"/opt/ml/output"</span>)) | |
| <span class="hljs-keyword">return</span> parser.parse_args() | |
| <span class="hljs-keyword">def</span> <span class="hljs-title function_">main</span>(): | |
| args = parse_args() | |
| <span class="hljs-comment"># the dataset downloads from the Hugging Face Hub inside the training container</span> | |
| dataset = load_dataset(args.dataset_name, split=<span class="hljs-string">"train"</span>) | |
| training_args = SFTConfig( | |
| output_dir=args.output_dir, | |
| max_steps=args.max_steps, | |
| per_device_train_batch_size=args.train_batch_size, | |
| learning_rate=args.learning_rate, | |
| logging_steps=<span class="hljs-number">5</span>, | |
| <span class="hljs-comment"># the final model is saved explicitly below</span> | |
| save_strategy=<span class="hljs-string">"no"</span>, | |
| report_to=[], | |
| ) | |
| trainer = SFTTrainer( | |
| model=args.model_name, | |
| args=training_args, | |
| train_dataset=dataset, | |
| ) | |
| trainer.train() | |
| <span class="hljs-comment"># save the model and tokenizer where SageMaker expects them</span> | |
| trainer.save_model(args.model_dir) | |
| trainer.processing_class.save_pretrained(args.model_dir) | |
| <span class="hljs-keyword">if</span> __name__ == <span class="hljs-string">"__main__"</span>: | |
| main()<!----></pre></div><!----> <blockquote class="note"><p>SageMaker does not support argparse actions. For example, if you want a boolean hyperparameter, specify <code>type</code> as <code>bool</code> in your script and provide an explicit <code>True</code> or <code>False</code> value.</p></blockquote> <!--[1--><h2 class="relative group"><a id="create-a-modeltrainer" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#create-a-modeltrainer"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Create a ModelTrainer</span></h2><!--]--><!----> <p>The <code>ModelTrainer</code> handles end-to-end SageMaker training. The most important parameters:</p> <ol><li><code>source_code</code> specifies the training script (<code>entry_script</code>) and its directory (<code>source_dir</code>).</li> <li><code>compute</code> specifies the instance(s) to launch. Refer to <a href="https://aws.amazon.com/sagemaker/pricing/" rel="nofollow">SageMaker pricing</a> for a complete list of instance types.</li> <li><code>training_image</code> is the training container image, retrieved with <code>image_uris.retrieve</code>.</li> <li><code>hyperparameters</code> are passed to the script as <code>--key value</code> command-line arguments.</li></ol> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-python "><!----><span class="hljs-keyword">from</span> sagemaker.core.helper.session_helper <span class="hljs-keyword">import</span> Session, get_execution_role | |
| <span class="hljs-keyword">from</span> sagemaker.train.model_trainer <span class="hljs-keyword">import</span> ModelTrainer | |
| <span class="hljs-keyword">from</span> sagemaker.train.configs <span class="hljs-keyword">import</span> SourceCode, Compute, StoppingCondition | |
| <span class="hljs-keyword">from</span> sagemaker.core <span class="hljs-keyword">import</span> image_uris | |
| <span class="hljs-comment"># set up the SageMaker session and execution role</span> | |
| sess = Session() | |
| role = get_execution_role() | |
| hyperparameters = { | |
| <span class="hljs-comment"># any small causal LM from the Hub works</span> | |
| <span class="hljs-string">"model_name"</span>: <span class="hljs-string">"Qwen/Qwen3-0.6B"</span>, | |
| <span class="hljs-comment"># conversational SFT dataset</span> | |
| <span class="hljs-string">"dataset_name"</span>: <span class="hljs-string">"trl-lib/Capybara"</span>, | |
| <span class="hljs-comment"># short run: enough to see the loss go down</span> | |
| <span class="hljs-string">"max_steps"</span>: <span class="hljs-number">50</span>, | |
| <span class="hljs-string">"train_batch_size"</span>: <span class="hljs-number">4</span>, | |
| <span class="hljs-string">"learning_rate"</span>: <span class="hljs-number">2e-5</span>, | |
| } | |
| instance_type = <span class="hljs-string">"ml.g6.xlarge"</span> | |
| <span class="hljs-comment"># Retrieve the Hugging Face PyTorch training DLC image URI</span> | |
| training_image = image_uris.retrieve( | |
| framework=<span class="hljs-string">"huggingface"</span>, | |
| region=sess.boto_region_name, | |
| <span class="hljs-comment"># Transformers version</span> | |
| version=<span class="hljs-string">"5.3.0"</span>, | |
| <span class="hljs-comment"># PyTorch version</span> | |
| base_framework_version=<span class="hljs-string">"pytorch2.9.0"</span>, | |
| <span class="hljs-comment"># Python version</span> | |
| py_version=<span class="hljs-string">"py312"</span>, | |
| image_scope=<span class="hljs-string">"training"</span>, | |
| instance_type=instance_type, | |
| ) | |
| model_trainer = ModelTrainer( | |
| sagemaker_session=sess, | |
| role=role, | |
| training_image=training_image, | |
| source_code=SourceCode( | |
| <span class="hljs-comment"># directory with the training script</span> | |
| source_dir=<span class="hljs-string">"./scripts"</span>, | |
| <span class="hljs-comment"># script to run in the training job</span> | |
| entry_script=<span class="hljs-string">"train.py"</span>, | |
| ), | |
| compute=Compute( | |
| instance_type=instance_type, | |
| instance_count=<span class="hljs-number">1</span>, | |
| <span class="hljs-comment"># uncomment for managed spot instances (needs spot quota)</span> | |
| <span class="hljs-comment"># enable_managed_spot_training=True,</span> | |
| ), | |
| stopping_condition=StoppingCondition( | |
| <span class="hljs-comment"># safety cap on billable seconds</span> | |
| max_runtime_in_seconds=<span class="hljs-number">3600</span>, | |
| ), | |
| hyperparameters=hyperparameters, | |
| )<!----></pre></div><!----> <p>If you are running a <code>TrainingJob</code> locally, define <code>instance_type='local'</code> or <code>instance_type='local_gpu'</code> for GPU usage. Note that this will not work with SageMaker Studio.</p> <p>The sections below reuse <code>sess</code>, <code>role</code>, <code>training_image</code>, and <code>hyperparameters</code> from this example; each snippet shows only what it changes.</p> <!--[1--><h2 class="relative group"><a id="start-the-training-job" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#start-the-training-job"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Start the training job</span></h2><!--]--><!----> <p>Call <code>train</code> to launch the job:</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-python "><!---->model_trainer.train()<!----></pre></div><!----> <p>SageMaker starts the instance, runs <code>train.py</code> with your hyperparameters, streams the logs, and uploads the model artifacts to S3 when the job finishes. The example script downloads its dataset from the Hub inside the container, so there is no data to upload.</p> <p>If your data lives in S3, pass it as input channels instead. Each channel is mounted inside the container at <code>/opt/ml/input/data/<channel_name></code> and exposed to your script as the <code>SM_CHANNEL_<channel_name></code> environment variable:</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-python "><!----><span class="hljs-keyword">from</span> sagemaker.train.configs <span class="hljs-keyword">import</span> InputData | |
| model_trainer.train( | |
| input_data_config=[ | |
| InputData(channel_name=<span class="hljs-string">"train"</span>, data_source=<span class="hljs-string">"s3://<your-bucket>/dataset/train"</span>), | |
| InputData(channel_name=<span class="hljs-string">"test"</span>, data_source=<span class="hljs-string">"s3://<your-bucket>/dataset/test"</span>), | |
| ] | |
| )<!----></pre></div><!----> <p>A channel <code>data_source</code> can be an S3 URI or a <code>FileSystemInput</code> for Amazon EFS or FSx for Lustre.</p> <!--[1--><h2 class="relative group"><a id="training-output-and-checkpoints" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#training-output-and-checkpoints"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Training output and checkpoints</span></h2><!--]--><!----> <p>If <code>output_dir</code> in the training arguments is set to <code>/opt/ml/model</code>, all training artifacts — logs, checkpoints, and models — are saved there. Amazon SageMaker archives the whole <code>/opt/ml/model</code> directory as <code>model.tar.gz</code> and uploads it to Amazon S3 at the end of the training job. Depending on your hyperparameters, this can lead to a large artifact (> 5GB), which slows down deployment for Amazon SageMaker Inference.</p> <p>You can control how checkpoints, logs, and artifacts are saved by customizing the training arguments. For example, set <code>save_total_limit</code> to cap the number of checkpoints: older checkpoints in <code>output_dir</code> are deleted once the limit is reached.</p> <p>To save artifacts continuously during training instead of only at the end, SageMaker supports <a href="https://docs.aws.amazon.com/sagemaker/latest/dg/model-checkpoints.html" rel="nofollow">checkpointing</a>: provide a <code>CheckpointConfig(s3_uri=...)</code> on the <code>ModelTrainer</code> and set <code>output_dir</code> to <code>/opt/ml/checkpoints</code>. In the example script, also switch <code>save_strategy</code> from <code>"no"</code> to <code>"steps"</code> so checkpoints are actually written.</p> <blockquote class="warning"><p>If you set <code>output_dir</code> to <code>/opt/ml/checkpoints</code>, call <code>trainer.save_model("/opt/ml/model")</code> — or <code>model.save_pretrained("/opt/ml/model")</code> and <code>tokenizer.save_pretrained("/opt/ml/model")</code> — at the end of training. Otherwise the model artifacts are missing from <code>model.tar.gz</code> and the model cannot be deployed to Amazon SageMaker for inference.</p></blockquote> <!--[1--><h2 class="relative group"><a id="access-the-trained-model" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#access-the-trained-model"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Access the trained model</span></h2><!--]--><!----> <p>Once training is complete, you can access your model through the <a href="https://console.aws.amazon.com/console/home?nc2=h_ct&src=header-signin" rel="nofollow">AWS console</a> or download it directly from S3. The S3 URI of the trained model artifacts is available on the completed training job:</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-python "><!----><span class="hljs-keyword">import</span> boto3 | |
| <span class="hljs-keyword">from</span> urllib.parse <span class="hljs-keyword">import</span> urlparse | |
| <span class="hljs-comment"># S3 URI where the trained model artifacts (model.tar.gz) are located</span> | |
| model_data = model_trainer._latest_training_job.model_artifacts.s3_model_artifacts | |
| parsed = urlparse(model_data) | |
| boto3.client(<span class="hljs-string">"s3"</span>).download_file( | |
| <span class="hljs-comment"># bucket</span> | |
| parsed.netloc, | |
| <span class="hljs-comment"># key</span> | |
| parsed.path.lstrip(<span class="hljs-string">"/"</span>), | |
| <span class="hljs-comment"># local path where the artifact is saved</span> | |
| <span class="hljs-string">"model.tar.gz"</span>, | |
| )<!----></pre></div><!----> <!--[1--><h2 class="relative group"><a id="distributed-training" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#distributed-training"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Distributed training</span></h2><!--]--><!----> <p>SageMaker provides two strategies for distributed training: data parallelism and model parallelism. Data parallelism splits a training set across several GPUs, while model parallelism splits a model across several GPUs.</p> <!--[2--><h3 class="relative group"><a id="data-parallelism" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#data-parallelism"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Data parallelism</span></h3><!--]--><!----> <p>The Hugging Face <code>Trainer</code> and the TRL trainers support distributed data parallel training. With <code>ModelTrainer</code> you launch your script with <code>torchrun</code> by passing a <code>Torchrun</code> config to the <code>distributed</code> parameter. Set <code>process_count_per_node</code> to the number of GPUs per instance (<code>ml.g6e.12xlarge</code> has 4):</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-python "><!----><span class="hljs-keyword">from</span> sagemaker.train.distributed <span class="hljs-keyword">import</span> Torchrun | |
| <span class="hljs-comment"># reuses sess, role, training_image, and hyperparameters from the ModelTrainer example above</span> | |
| <span class="hljs-comment"># 4x L40S GPUs</span> | |
| instance_type = <span class="hljs-string">"ml.g6e.12xlarge"</span> | |
| <span class="hljs-comment"># create the ModelTrainer with torchrun for distributed data parallelism</span> | |
| model_trainer = ModelTrainer( | |
| sagemaker_session=sess, | |
| role=role, | |
| training_image=training_image, | |
| source_code=SourceCode(source_dir=<span class="hljs-string">"./scripts"</span>, entry_script=<span class="hljs-string">"train.py"</span>), | |
| compute=Compute(instance_type=instance_type, instance_count=<span class="hljs-number">2</span>), | |
| distributed=Torchrun(process_count_per_node=<span class="hljs-number">4</span>), | |
| hyperparameters=hyperparameters, | |
| )<!----></pre></div><!----> <!--[2--><h3 class="relative group"><a id="model-parallelism" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#model-parallelism"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Model parallelism</span></h3><!--]--><!----> <p>For models too large for a single GPU, the SageMaker Model Parallelism library (SMP) provides tensor parallelism, context parallelism, and sharded data parallelism. Enable it by passing an <code>SMP</code> config to <code>Torchrun</code>:</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-python "><!----><span class="hljs-keyword">from</span> sagemaker.train.distributed <span class="hljs-keyword">import</span> Torchrun, SMP | |
| <span class="hljs-comment"># reuses sess, role, training_image, and hyperparameters from the ModelTrainer example above</span> | |
| <span class="hljs-comment"># 8x A100 GPUs</span> | |
| instance_type = <span class="hljs-string">"ml.p4de.24xlarge"</span> | |
| <span class="hljs-comment"># create the ModelTrainer with torchrun + SMP for model parallelism</span> | |
| model_trainer = ModelTrainer( | |
| sagemaker_session=sess, | |
| role=role, | |
| training_image=training_image, | |
| source_code=SourceCode(source_dir=<span class="hljs-string">"./scripts"</span>, entry_script=<span class="hljs-string">"train.py"</span>), | |
| compute=Compute(instance_type=instance_type, instance_count=<span class="hljs-number">2</span>), | |
| distributed=Torchrun( | |
| process_count_per_node=<span class="hljs-number">8</span>, | |
| smp=SMP( | |
| tensor_parallel_degree=<span class="hljs-number">2</span>, | |
| hybrid_shard_degree=<span class="hljs-number">1</span>, | |
| ), | |
| ), | |
| hyperparameters=hyperparameters, | |
| )<!----></pre></div><!----> <!--[1--><h2 class="relative group"><a id="spot-instances" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#spot-instances"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Spot instances</span></h2><!--]--><!----> <p>Managed spot training uses <a href="https://docs.aws.amazon.com/sagemaker/latest/dg/model-managed-spot-training.html" rel="nofollow">fully-managed EC2 spot instances</a> and can save up to 90% of training costs. Set <code>enable_managed_spot_training=True</code> on <code>Compute</code>, define <code>max_wait_time_in_seconds</code> and <code>max_runtime_in_seconds</code> on <code>StoppingCondition</code>, and enable checkpointing so an interrupted job can resume:</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-python "><!----><span class="hljs-keyword">from</span> sagemaker.train.configs <span class="hljs-keyword">import</span> StoppingCondition, CheckpointConfig | |
| <span class="hljs-comment"># reuses sess, role, and training_image from the ModelTrainer example above</span> | |
| <span class="hljs-comment"># spot jobs can be interrupted, so the script must write checkpoints to /opt/ml/checkpoints</span> | |
| hyperparameters = { | |
| <span class="hljs-string">"model_name"</span>: <span class="hljs-string">"Qwen/Qwen3-0.6B"</span>, | |
| <span class="hljs-string">"dataset_name"</span>: <span class="hljs-string">"trl-lib/Capybara"</span>, | |
| <span class="hljs-string">"max_steps"</span>: <span class="hljs-number">50</span>, | |
| <span class="hljs-string">"train_batch_size"</span>: <span class="hljs-number">4</span>, | |
| <span class="hljs-string">"learning_rate"</span>: <span class="hljs-number">2e-5</span>, | |
| <span class="hljs-string">"output_dir"</span>: <span class="hljs-string">"/opt/ml/checkpoints"</span>, | |
| } | |
| model_trainer = ModelTrainer( | |
| sagemaker_session=sess, | |
| role=role, | |
| training_image=training_image, | |
| source_code=SourceCode(source_dir=<span class="hljs-string">"./scripts"</span>, entry_script=<span class="hljs-string">"train.py"</span>), | |
| compute=Compute( | |
| instance_type=<span class="hljs-string">"ml.g6.xlarge"</span>, | |
| instance_count=<span class="hljs-number">1</span>, | |
| <span class="hljs-comment"># use fully-managed spot instances</span> | |
| enable_managed_spot_training=<span class="hljs-literal">True</span>, | |
| ), | |
| <span class="hljs-comment"># max_wait_time_in_seconds should be equal to or greater than max_runtime_in_seconds</span> | |
| stopping_condition=StoppingCondition( | |
| max_runtime_in_seconds=<span class="hljs-number">3600</span>, | |
| max_wait_time_in_seconds=<span class="hljs-number">7200</span>, | |
| ), | |
| checkpoint_config=CheckpointConfig(s3_uri=<span class="hljs-string">f"s3://<span class="hljs-subst">{sess.default_bucket()}</span>/checkpoints"</span>), | |
| hyperparameters=hyperparameters, | |
| )<!----></pre></div><!----> <blockquote class="note"><p>Spot and on-demand quotas are separate, and new accounts can start with a spot limit of 0. If job creation fails with <code>ResourceLimitExceeded</code>, check your <a href="https://console.aws.amazon.com/servicequotas/home/services/sagemaker/quotas" rel="nofollow">SageMaker quotas</a> or run on-demand.</p></blockquote> <!--[1--><h2 class="relative group"><a id="git-repository" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#git-repository"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>Git repository</span></h2><!--]--><!----> <p>The v2 <code>git_config</code> parameter is not available in <code>ModelTrainer</code>. To run a training script that lives in a GitHub repository (such as the <a href="https://github.com/huggingface/transformers/tree/main/examples" rel="nofollow">🤗 Transformers example scripts</a>), clone the repository locally first and point <code>source_dir</code>/<code>entry_script</code> at the checked-out files. Choose a branch that matches the Transformers version of your training image.</p> <blockquote class="tip"><p>Save your model to S3 by setting <code>output_dir=/opt/ml/model</code> in the hyperparameters of your training script.</p></blockquote> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-bash "><!----><span class="hljs-comment"># clone the repo locally, matching the transformers version of your training image</span> | |
| git <span class="hljs-built_in">clone</span> --branch v5.3.0 https://github.com/huggingface/transformers.git<!----></pre></div><!----> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-python "><!----><span class="hljs-comment"># reuses sess, role, and training_image from the ModelTrainer example above</span> | |
| <span class="hljs-comment"># run_glue.py takes the Transformers example argument names</span> | |
| hyperparameters = { | |
| <span class="hljs-string">"epochs"</span>: <span class="hljs-number">1</span>, | |
| <span class="hljs-string">"per_device_train_batch_size"</span>: <span class="hljs-number">32</span>, | |
| <span class="hljs-string">"model_name_or_path"</span>: <span class="hljs-string">"distilbert-base-uncased"</span>, | |
| } | |
| <span class="hljs-comment"># create the ModelTrainer pointing at the cloned example directory</span> | |
| model_trainer = ModelTrainer( | |
| sagemaker_session=sess, | |
| role=role, | |
| training_image=training_image, | |
| source_code=SourceCode( | |
| source_dir=<span class="hljs-string">"transformers/examples/pytorch/text-classification"</span>, | |
| entry_script=<span class="hljs-string">"run_glue.py"</span>, | |
| requirements=<span class="hljs-string">"requirements.txt"</span>, | |
| ), | |
| compute=Compute(instance_type=<span class="hljs-string">"ml.g6.xlarge"</span>, instance_count=<span class="hljs-number">1</span>), | |
| hyperparameters=hyperparameters, | |
| )<!----></pre></div><!----> <!--[1--><h2 class="relative group"><a id="sagemaker-metrics" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#sagemaker-metrics"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>SageMaker metrics</span></h2><!--]--><!----> <p><a href="https://docs.aws.amazon.com/sagemaker/latest/dg/training-metrics.html#define-train-metrics" rel="nofollow">SageMaker metrics</a> automatically parse training job logs and send metrics to CloudWatch. Specify each metric’s name and a regular expression for SageMaker to match. With <code>ModelTrainer</code> you attach them using <code>with_metric_definitions</code>:</p> <div class="code-block relative "><div class="absolute top-2.5 right-4"><button class="inline-flex items-center relative text-sm focus:text-green-500 cursor-pointer focus:outline-none transition duration-200 ease-in-out opacity-0 mx-0.5 text-gray-600 " title="code excerpt" type="button"><svg xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M28,10V28H10V10H28m0-2H10a2,2,0,0,0-2,2V28a2,2,0,0,0,2,2H28a2,2,0,0,0,2-2V10a2,2,0,0,0-2-2Z" transform="translate(0)"></path><path d="M4,18H2V4A2,2,0,0,1,4,2H18V4H4Z" transform="translate(0)"></path><rect fill="none" width="32" height="32"></rect></svg><!----> <div class=" absolute pointer-events-none transition-opacity bg-black text-white py-1 px-2 leading-tight rounded font-normal shadow left-1/2 top-full transform -translate-x-1/2 translate-y-2 opacity-0 "><div class="absolute bottom-full left-1/2 transform -translate-x-1/2 w-0 h-0 border-black border-4 border-t-0" style="border-left-color: transparent; border-right-color: transparent;"></div> Copied</div><!----></button><!----></div> <pre class="language-python "><!----><span class="hljs-keyword">from</span> sagemaker.train.configs <span class="hljs-keyword">import</span> MetricDefinition | |
| <span class="hljs-comment"># reuses sess, role, training_image, and hyperparameters from the ModelTrainer example above</span> | |
| <span class="hljs-comment"># SFTTrainer logs lines like {'loss': 2.34, ...}; parse the loss into CloudWatch</span> | |
| metric_definitions = [ | |
| MetricDefinition(name=<span class="hljs-string">"train-loss"</span>, regex=<span class="hljs-string">"'loss': ([0-9.]+)"</span>), | |
| ] | |
| model_trainer = ModelTrainer( | |
| sagemaker_session=sess, | |
| role=role, | |
| training_image=training_image, | |
| source_code=SourceCode(source_dir=<span class="hljs-string">"./scripts"</span>, entry_script=<span class="hljs-string">"train.py"</span>), | |
| compute=Compute(instance_type=<span class="hljs-string">"ml.g6.xlarge"</span>, instance_count=<span class="hljs-number">1</span>), | |
| hyperparameters=hyperparameters, | |
| ).with_metric_definitions(metric_definitions)<!----></pre></div><!----> <!--[1--><h2 class="relative group"><a id="whats-next" class="header-link block pr-1.5 text-lg no-hover:hidden with-hover:absolute with-hover:p-1.5 with-hover:opacity-0 with-hover:group-hover:opacity-100 with-hover:right-full" href="#whats-next"><span><svg xmlns="http://www.w3.org/2000/svg" xmlns:xlink="http://www.w3.org/1999/xlink" aria-hidden="true" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 256 256"><path d="M167.594 88.393a8.001 8.001 0 0 1 0 11.314l-67.882 67.882a8 8 0 1 1-11.314-11.315l67.882-67.881a8.003 8.003 0 0 1 11.314 0zm-28.287 84.86l-28.284 28.284a40 40 0 0 1-56.567-56.567l28.284-28.284a8 8 0 0 0-11.315-11.315l-28.284 28.284a56 56 0 0 0 79.196 79.197l28.285-28.285a8 8 0 1 0-11.315-11.314zM212.852 43.14a56.002 56.002 0 0 0-79.196 0l-28.284 28.284a8 8 0 1 0 11.314 11.314l28.284-28.284a40 40 0 0 1 56.568 56.567l-28.285 28.285a8 8 0 0 0 11.315 11.314l28.284-28.284a56.065 56.065 0 0 0 0-79.196z" fill="currentColor"></path></svg><!----></span></a> <span>What’s next</span></h2><!--]--><!----> <p>Once your training job is complete, the model artifacts are in S3 and ready for deployment. Continue with <a href="./deploy-sagemaker-sdk#deploy-a--transformers-model-trained-in-sagemaker">Deploy models</a> to serve your trained model on a SageMaker endpoint.</p> <a class="!text-gray-400 !no-underline text-sm flex items-center not-prose mt-4" href="https://github.com/huggingface/hub-docs/blob/main/docs/sagemaker/source/tutorials/sagemaker-sdk/training-sagemaker-sdk.md" target="_blank"><svg class="mr-1" xmlns="http://www.w3.org/2000/svg" aria-hidden="true" fill="currentColor" focusable="false" role="img" width="1em" height="1em" preserveAspectRatio="xMidYMid meet" viewBox="0 0 32 32"><path d="M31,16l-7,7l-1.41-1.41L28.17,16l-5.58-5.59L24,9l7,7z"></path><path d="M1,16l7-7l1.41,1.41L3.83,16l5.58,5.59L8,23l-7-7z"></path><path d="M12.419,25.484L17.639,6.552l1.932,0.518L14.351,26.002z"></path></svg><!----> <span><span class="underline">Update</span> on GitHub</span></a><!----> <p></p><!--]--><!----><!--]--><!--]--><!--]--> <!--[-1--><!--]--><!--]--> | |
| <script> | |
| { | |
| __sveltekit_dwws3f = { | |
| base: "/docs/sagemaker/pr_2709/en", | |
| assets: "/docs/sagemaker/pr_2709/en" | |
| }; | |
| const element = document.currentScript.parentElement; | |
| Promise.all([ | |
| import("/docs/sagemaker/pr_2709/en/_app/immutable/entry/start.COiK-bUV.js"), | |
| import("/docs/sagemaker/pr_2709/en/_app/immutable/entry/app.CjYANwwM.js") | |
| ]).then(([kit, app]) => { | |
| kit.start(app, element, { | |
| node_ids: [0, 30], | |
| data: [null,null], | |
| form: null, | |
| error: null | |
| }); | |
| }); | |
| } | |
| </script> | |
Xet Storage Details
- Size:
- 61.2 kB
- Xet hash:
- 6dc6dc3b0fb66555a18b0549564a9fee34fef71d95affd73caca11e82e8b27bf
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.