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
4 changes: 3 additions & 1 deletion sagemaker-train/src/sagemaker/train/tuner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
29 changes: 29 additions & 0 deletions sagemaker-train/tests/unit/train/test_tuner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading