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
85 changes: 73 additions & 12 deletions sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -643,23 +643,78 @@ def _validate_networking(vpc):
"VPC config requires both non-empty 'security_group_ids' and 'subnets'."
)

def _nova_managed_configuration(self):
"""Return the ManagedConfiguration required for the base model, or None.

Nova (closed-source) models require every ModelPackageGroup used by the
job — output and intermediate checkpoint — to use Restricted managed
storage. Open-source models have no managed-storage requirement.
"""
if not _is_nova_model(self._model_name):
return None
from sagemaker.core.shapes import ManagedConfiguration

return ManagedConfiguration(managed_storage_type="Restricted")

def _validate_mpg_managed_storage(self, mpg, managed_configuration) -> None:
"""Validate an existing ModelPackageGroup against a required ManagedConfiguration.

Raises:
ValueError: If ``managed_configuration`` is required and the existing
group's managed storage type does not match. Existing groups
cannot be converted, so failing fast here (at trainer
construction) is preferable to a late job-submission rejection.
"""
if managed_configuration is None:
return
required = managed_configuration.managed_storage_type
existing = getattr(mpg, "managed_configuration", None)
existing_type = getattr(existing, "managed_storage_type", None)
if not isinstance(existing_type, str):
existing_type = None
if existing_type != required:
group_name = getattr(mpg, "model_package_group_name", None) or getattr(
mpg, "model_package_group_arn", "<unknown>"
)
raise ValueError(
f"ModelPackageGroup '{group_name}' uses "
f"'{existing_type or 'Standard'}' managed storage, but model "
f"'{self._model_name}' requires '{required}'. Existing groups cannot "
"be converted. Pass a ModelPackageGroup created with "
f"ManagedConfiguration(managed_storage_type='{required}'), or omit "
"the parameter to auto-create a compliant group."
)

def _get_or_create_mpg(self, value, default_name: str, session, managed_configuration=None) -> str:
"""Resolve an existing ModelPackageGroup or auto-create one.

If ``value`` is provided (object or string), validates it exists and returns its ARN.
If ``value`` is None, creates a ModelPackageGroup with ``default_name`` (get-or-create).
When ``managed_configuration`` is provided, every resolution path validates
that the resulting group satisfies it (auto-created groups are created
with it).

Returns:
The ModelPackageGroup ARN.
"""
if value:
if isinstance(value, ModelPackageGroup):
if managed_configuration is not None:
# A caller-constructed object may not carry managed_configuration;
# fetch the authoritative record before validating.
fetched = ModelPackageGroup.get(
model_package_group_name=value.model_package_group_name,
session=session.boto_session,
region=session.boto_session.region_name,
)
self._validate_mpg_managed_storage(fetched, managed_configuration)
return value.model_package_group_arn
mpg = ModelPackageGroup.get(
model_package_group_name=value,
session=session.boto_session,
region=session.boto_session.region_name,
)
self._validate_mpg_managed_storage(mpg, managed_configuration)
return mpg.model_package_group_arn

# Auto-create (get-or-create with deterministic name)
Expand All @@ -670,6 +725,9 @@ def _get_or_create_mpg(self, value, default_name: str, session, managed_configur
session=session.boto_session,
region=session.boto_session.region_name,
)
self._validate_mpg_managed_storage(mpg, managed_configuration)
except ValueError:
raise
except Exception:
try:
create_kwargs = {
Expand All @@ -695,22 +753,27 @@ def _resolve_model_package_group(self, model, output_model_package_group, sessio
2. If ``model`` is a ModelPackage, derives the group from it.
3. Otherwise, auto-creates ``{model_name}-mtrl-mpg`` (get-or-create).

For Nova models, every branch validates (or creates) the group with
Restricted managed storage — including the explicit and
continued-customization (ModelPackage-derived) branches.

Returns:
The ModelPackageGroup ARN.
"""
managed_config = self._nova_managed_configuration()

if output_model_package_group:
return self._get_or_create_mpg(output_model_package_group, None, session)
return self._get_or_create_mpg(
output_model_package_group, None, session, managed_configuration=managed_config
)

# Derive from ModelPackage
if isinstance(model, ModelPackage):
group_name = model.model_package_group_name
if group_name:
return self._get_or_create_mpg(group_name, None, session)

managed_config = None
if _is_nova_model(self._model_name):
from sagemaker.core.shapes import ManagedConfiguration
managed_config = ManagedConfiguration(managed_storage_type="Restricted")
return self._get_or_create_mpg(
group_name, None, session, managed_configuration=managed_config
)

return self._get_or_create_mpg(
None, f"{self._model_name}-mtrl-mpg", session, managed_configuration=managed_config
Expand All @@ -719,17 +782,15 @@ def _resolve_model_package_group(self, model, output_model_package_group, sessio
def _resolve_intermediate_checkpoint_mpg(self, intermediate_checkpoint_mpg, session) -> str:
"""Resolve or auto-create the intermediate checkpoint ModelPackageGroup.

If provided, validates it exists. Otherwise auto-creates
If provided, validates it exists (and, for Nova models, that it uses
Restricted managed storage). Otherwise auto-creates
``{model_name}-mtrl-checkpoint-mpg`` (get-or-create).
Raises ValueError if the resolved ARN is the same as ``output_model_package_group``.

Returns:
The ModelPackageGroup ARN.
"""
managed_config = None
if not intermediate_checkpoint_mpg and _is_nova_model(self._model_name):
from sagemaker.core.shapes import ManagedConfiguration
managed_config = ManagedConfiguration(managed_storage_type="Restricted")
managed_config = self._nova_managed_configuration()

arn = self._get_or_create_mpg(
intermediate_checkpoint_mpg,
Expand Down
103 changes: 103 additions & 0 deletions sagemaker-train/tests/unit/train/test_multi_turn_rl_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -491,6 +491,109 @@ def test_nova_model_creates_restricted_mpg(self, mock_get, mock_create):
call_kwargs = mock_create.call_args[1]
assert call_kwargs["managed_configuration"].managed_storage_type == "Restricted"

@staticmethod
def _mock_mpg(storage_type=None, name="my-group", arn=MPG_ARN):
mpg = MagicMock()
mpg.model_package_group_name = name
mpg.model_package_group_arn = arn
if storage_type is None:
mpg.managed_configuration = None
else:
mpg.managed_configuration = MagicMock()
mpg.managed_configuration.managed_storage_type = storage_type
return mpg

@patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get")
def test_nova_explicit_standard_group_raises(self, mock_get):
"""Nova + explicit output group with Standard storage must fail fast."""
mock_get.return_value = self._mock_mpg(storage_type=None)

trainer = self._make_trainer()
trainer._model_name = "amazon-nova-pro"
with pytest.raises(ValueError, match="requires 'Restricted'"):
trainer._resolve_model_package_group("amazon-nova-pro", "my-group", self._mock_session())

@patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get")
def test_nova_explicit_restricted_group_passes(self, mock_get):
mock_get.return_value = self._mock_mpg(storage_type="Restricted")

trainer = self._make_trainer()
trainer._model_name = "amazon-nova-pro"
result = trainer._resolve_model_package_group(
"amazon-nova-pro", "my-group", self._mock_session()
)
assert result == MPG_ARN

@patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get")
def test_nova_mpg_object_standard_group_raises(self, mock_get):
"""Nova + explicit ModelPackageGroup object is validated via a fresh fetch."""
from sagemaker.core.resources import ModelPackageGroup as MPG

mock_get.return_value = self._mock_mpg(storage_type=None)
mpg_obj = MagicMock(spec=MPG)
mpg_obj.model_package_group_name = "my-group"
mpg_obj.model_package_group_arn = MPG_ARN

trainer = self._make_trainer()
trainer._model_name = "amazon-nova-pro"
with pytest.raises(ValueError, match="requires 'Restricted'"):
trainer._resolve_model_package_group("amazon-nova-pro", mpg_obj, self._mock_session())

@patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get")
def test_nova_derived_standard_group_raises(self, mock_get):
"""Nova continued-customization: group derived from source ModelPackage must be Restricted."""
mock_get.return_value = self._mock_mpg(storage_type=None, name="derived-group")
mock_model = MagicMock(spec=ModelPackage)
mock_model.model_package_group_name = "derived-group"

trainer = self._make_trainer()
trainer._model_name = "amazon-nova-pro"
with pytest.raises(ValueError, match="requires 'Restricted'"):
trainer._resolve_model_package_group(mock_model, None, self._mock_session())

@patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get")
def test_nova_derived_restricted_group_passes(self, mock_get):
mock_get.return_value = self._mock_mpg(storage_type="Restricted", name="derived-group")
mock_model = MagicMock(spec=ModelPackage)
mock_model.model_package_group_name = "derived-group"

trainer = self._make_trainer()
trainer._model_name = "amazon-nova-pro"
result = trainer._resolve_model_package_group(mock_model, None, self._mock_session())
assert result == MPG_ARN

@patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get")
def test_nova_existing_default_named_standard_group_raises(self, mock_get):
"""Nova auto-create path: a pre-existing default-named Standard group must not be reused."""
mock_get.return_value = self._mock_mpg(
storage_type=None, name="amazon-nova-pro-mtrl-mpg"
)

trainer = self._make_trainer()
trainer._model_name = "amazon-nova-pro"
with pytest.raises(ValueError, match="requires 'Restricted'"):
trainer._resolve_model_package_group("amazon-nova-pro", None, self._mock_session())

@patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get")
def test_oss_model_standard_group_passes_unvalidated(self, mock_get):
"""Non-Nova models keep the existing behavior: no storage validation."""
mock_get.return_value = self._mock_mpg(storage_type=None)

trainer = self._make_trainer()
result = trainer._resolve_model_package_group("test-model", "my-group", self._mock_session())
assert result == MPG_ARN

@patch("sagemaker.train.multi_turn_rl_trainer.ModelPackageGroup.get")
def test_nova_intermediate_explicit_standard_group_raises(self, mock_get):
"""Nova + explicit intermediate checkpoint group with Standard storage must fail."""
mock_get.return_value = self._mock_mpg(storage_type=None, name="my-ckpt-group")

trainer = self._make_trainer()
trainer._model_name = "amazon-nova-pro"
trainer.output_model_package_group = "arn:other"
with pytest.raises(ValueError, match="requires 'Restricted'"):
trainer._resolve_intermediate_checkpoint_mpg("my-ckpt-group", self._mock_session())


class TestAgentRuntimeIdPattern:
def test_valid_runtime_id(self):
Expand Down
Loading