From 6db012fdd8bcde4ab3ab2f22add7d815167ae361 Mon Sep 17 00:00:00 2001 From: Mohamed Zeidan Date: Thu, 24 Sep 2026 14:56:07 -0700 Subject: [PATCH] fix: preserve content_type when HyperparameterTuner converts InputData to Channel (#5632) --- sagemaker-train/src/sagemaker/train/tuner.py | 4 ++- .../tests/unit/train/test_tuner.py | 29 +++++++++++++++++++ 2 files changed, 32 insertions(+), 1 deletion(-) diff --git a/sagemaker-train/src/sagemaker/train/tuner.py b/sagemaker-train/src/sagemaker/train/tuner.py index ed872dc894..c11cc60740 100644 --- a/sagemaker-train/src/sagemaker/train/tuner.py +++ b/sagemaker-train/src/sagemaker/train/tuner.py @@ -1424,10 +1424,12 @@ def _build_training_job_definition(self, inputs): # List of InputData or Channel objects for inp in inputs: if isinstance(inp, InputData): - # Convert InputData to Channel + # Convert InputData to Channel. Preserve content_type so built-in + # algorithms (e.g. XGBoost) know the data format (issue #5632). input_data_config.append( Channel( channel_name=inp.channel_name, + content_type=inp.content_type, data_source=DataSource( s3_data_source=S3DataSource( s3_data_type="S3Prefix", diff --git a/sagemaker-train/tests/unit/train/test_tuner.py b/sagemaker-train/tests/unit/train/test_tuner.py index d8010fa2d0..b3354fec35 100644 --- a/sagemaker-train/tests/unit/train/test_tuner.py +++ b/sagemaker-train/tests/unit/train/test_tuner.py @@ -581,6 +581,35 @@ def test_build_training_job_definition_includes_internal_channels(self): assert "validation" in channel_names, "User 'validation' channel should be included" assert len(channel_names) == 4, "Should have exactly 4 channels" + def test_build_training_job_definition_preserves_content_type(self): + """Regression for #5632. + + Converting an InputData to a Channel must carry over content_type, otherwise built-in + algorithms fail because the container doesn't know the data format. + """ + from sagemaker.core.training.configs import InputData + + tuner = HyperparameterTuner( + model_trainer=_create_mock_model_trainer(), + objective_metric_name="validation:auc", + hyperparameter_ranges=_create_single_hp_range(), + ) + + definition = tuner._build_training_job_definition( + [ + InputData( + channel_name="train", + data_source="s3://bucket/train/train.csv", + content_type="csv", + ) + ] + ) + + train_channel = next( + ch for ch in definition.input_data_config if ch.channel_name == "train" + ) + assert train_channel.content_type == "csv" + def test_build_training_job_definition_includes_spot_params(self): """Test that _build_training_job_definition includes spot parameters.""" tuner = HyperparameterTuner(