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: 4 additions & 0 deletions sagemaker-train/src/sagemaker/ai_registry/air_hub.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,7 @@ def import_hub_content(
hub_content_document: str,
hub_content_version: str = AIR_HUB_CONTENT_DEFAULT_VERSION,
tags: Optional[tuple] = None,
description: Optional[str] = None,
session: Optional[Session] = None,
):
"""Import hub content into the AI Registry hub.
Expand All @@ -108,6 +109,7 @@ def import_hub_content(
document_schema_version: Schema version of the document
hub_content_document: JSON document content
tags: Optional tuple of tags
description: Optional description of the hub content
session: Boto3 session

Returns:
Expand All @@ -127,6 +129,8 @@ def import_hub_content(
}
if tags:
request["HubContentSearchKeywords"] = [f"{tag[0]}:{tag[1]}" for tag in tags]
if description:
request["HubContentDescription"] = description
return client.import_hub_content(**request)

@classmethod
Expand Down
26 changes: 19 additions & 7 deletions sagemaker-train/src/sagemaker/ai_registry/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
from __future__ import annotations

import json
import logging
import os
import tempfile
from datetime import datetime
Expand Down Expand Up @@ -55,6 +56,8 @@
from sagemaker.core.helper.session_helper import Session
from sagemaker.train.defaults import TrainDefaults

logger = logging.getLogger(__name__)


class DataSet(AIRHubEntity):
"""Dataset entity for AI Registry."""
Expand Down Expand Up @@ -358,6 +361,7 @@ def create(
document_schema_version=DATASET_DOCUMENT_SCHEMA_VERSION,
hub_content_document=document_str,
tags=tags,
description=description,
session=sagemaker_session
)

Expand Down Expand Up @@ -519,15 +523,15 @@ def create_version(
self,
source: str,
customization_technique: Optional[CustomizationTechnique] = None
) -> bool:
) -> Optional["DataSet"]:
"""Create a new version of this dataset.

Args:
source: S3 URI or local file path for the dataset
customization_technique: Customization technique to use. If None, uses existing technique.

Returns:
True if version created successfully, False otherwise
DataSet: The newly created version, or None if creation failed
"""
try:
# Get current dataset metadata
Expand All @@ -545,19 +549,27 @@ def create_version(
technique = customization_technique or (CustomizationTechnique(existing_technique) if existing_technique else None)

# Create new version
DataSet.create(
new_dataset = DataSet.create(
name=self.name,
source=source,
customization_technique=technique,
tags=[
(TAG_KEY_CUSTOMIZATION_TECHNIQUE, technique.value),
(TAG_KEY_METHOD, keywords.get(TAG_KEY_METHOD, ""))
] if technique else [(TAG_KEY_METHOD, keywords.get(TAG_KEY_METHOD, ""))]
] if technique else [(TAG_KEY_METHOD, keywords.get(TAG_KEY_METHOD, ""))],
sagemaker_session=self.sagemaker_session,
)
logger.info(
"Created new version %s for dataset %s, arn: %s",
new_dataset.version, self.name, new_dataset.arn
)
return True
return new_dataset
except Exception as e:
print(f"Failed to create new version for dataset {self.name} with exception : {e}")
return False
logger.error(
"Failed to create new version for dataset %s with exception : %s",
self.name, e
)
return None

@staticmethod
def _parse_keywords(search_keywords: List[str]) -> dict:
Expand Down
31 changes: 26 additions & 5 deletions sagemaker-train/src/sagemaker/ai_registry/evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@

import io
import json
import logging
import os
import zipfile
from collections.abc import Sequence
Expand All @@ -40,6 +41,7 @@
RESPONSE_KEY_LAST_MODIFIED_TIME, RESPONSE_KEY_FUNCTION_ARN,
RESPONSE_KEY_HUB_CONTENT_NAME, RESPONSE_KEY_HUB_CONTENT_STATUS,
RESPONSE_KEY_HUB_CONTENT_DOCUMENT, RESPONSE_KEY_HUB_CONTENT_SEARCH_KEYWORDS,
RESPONSE_KEY_HUB_CONTENT_DESCRIPTION,
DOC_KEY_JSON_CONTENT,
DOC_KEY_REFERENCE, DOC_KEY_SUB_TYPE, REWARD_FUNCTION, REWARD_PROMPT,
)
Expand All @@ -52,6 +54,9 @@
from sagemaker.train.common_utils.finetune_utils import _get_current_domain_id
from sagemaker.train.defaults import TrainDefaults

logger = logging.getLogger(__name__)


class EvaluatorMethod(Enum):
"""Enum for Evaluator method types."""
BYOC = "byoc"
Expand Down Expand Up @@ -103,6 +108,7 @@ def __init__(
status: Optional[HubContentStatus] = None,
created_time: Optional[datetime] = None,
updated_time: Optional[datetime] = None,
description: Optional[str] = None,
sagemaker_session: Optional[Session] = None
) -> None:
"""Initialize Evaluator entity.
Expand All @@ -117,9 +123,10 @@ def __init__(
status: Current status of the evaluator
created_time: Creation timestamp
updated_time: Last update timestamp
description: Description of the evaluator
sagemaker_session: Optional SageMaker session.
"""
super().__init__(name, version, arn, status, created_time, updated_time,sagemaker_session)
super().__init__(name, version, arn, status, created_time, updated_time, description, sagemaker_session)
self.method = method
self.type = type
self.reference = reference
Expand Down Expand Up @@ -160,6 +167,7 @@ def refresh(self):
self.reference = json_content.get(DOC_KEY_REFERENCE, "")
self.type = json_content.get(DOC_KEY_SUB_TYPE, "")
self.status = response[RESPONSE_KEY_HUB_CONTENT_STATUS]
self.description = response.get(RESPONSE_KEY_HUB_CONTENT_DESCRIPTION, "")
method_str = keywords.get(TAG_KEY_METHOD)
self.method = EvaluatorMethod(method_str) if method_str else None
self.created = response.get(RESPONSE_KEY_CREATION_TIME)
Expand Down Expand Up @@ -198,6 +206,7 @@ def get(cls, name: str, sagemaker_session=None) -> "Evaluator":
method=EvaluatorMethod(method_str) if method_str else None,
reference=reference,
status=response[RESPONSE_KEY_HUB_CONTENT_STATUS],
description=response.get(RESPONSE_KEY_HUB_CONTENT_DESCRIPTION, ""),
created_time=response.get(RESPONSE_KEY_CREATION_TIME),
updated_time=response.get(RESPONSE_KEY_LAST_MODIFIED_TIME),
)
Expand All @@ -212,6 +221,7 @@ def create(
wait: bool = True,
role: Optional[str] = None,
domain_id: Optional[str] = None,
description: Optional[str] = None,
sagemaker_session: Optional[Session] = None,
) -> "Evaluator":
"""Create a new Evaluator entity in the AI Registry.
Expand All @@ -226,6 +236,7 @@ def create(
is visible in Studio. If not provided, it is auto-detected from the Studio
environment; supply it explicitly when creating evaluators outside Studio
(e.g. from a laptop or CI) so they still appear in the target domain.
description: Optional description of the evaluator.

Returns:
Evaluator: Newly created Evaluator instance
Expand Down Expand Up @@ -294,6 +305,7 @@ def create(
document_schema_version=EVALUATOR_DOCUMENT_SCHEMA_VERSION,
hub_content_document=hub_content_document,
tags=tags,
description=description,
session=sagemaker_session,
)

Expand All @@ -310,6 +322,7 @@ def create(
created_time=describe_response[RESPONSE_KEY_CREATION_TIME],
updated_time=describe_response[RESPONSE_KEY_LAST_MODIFIED_TIME],
reference=reference,
description=description,
sagemaker_session=sagemaker_session
)

Expand Down Expand Up @@ -495,21 +508,29 @@ def get_versions(self) -> List["Evaluator"]:
return evaluators

@_telemetry_emitter(feature=Feature.MODEL_CUSTOMIZATION, func_name="Evaluator.create_version")
def create_version(self, source: str) -> bool:
def create_version(self, source: str) -> "Evaluator":
"""Create a new version of this evaluator.

Args:
source: Lambda ARN or local file path for the function

Returns:
bool: True if version created successfully, False otherwise
Evaluator: The newly created version.

Raises:
RuntimeError: If version creation fails.
"""
try:
Evaluator.create(
new_evaluator = Evaluator.create(
name=self.name,
type=self.type,
source=source,
sagemaker_session=self.sagemaker_session,
)
logger.info(
"Created new version %s for evaluator %s, arn: %s",
new_evaluator.version, self.name, new_evaluator.arn
)
return True
return new_evaluator
except Exception as e:
raise RuntimeError(f"[PySDK Error] Failed to create new version: {str(e)}")
5 changes: 3 additions & 2 deletions sagemaker-train/tests/integ/ai_registry/test_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -184,9 +184,10 @@ def test_dataset_wait(self, unique_name, sample_jsonl_file, cleanup_list):
def test_create_dataset_version(self, unique_name, sample_jsonl_file, cleanup_list):
"""Test creating new dataset version."""
dataset = DataSet.create(name=unique_name, source=sample_jsonl_file, wait=False)
result = dataset.create_version(sample_jsonl_file)
new_version = dataset.create_version(sample_jsonl_file)
cleanup_list.append(dataset)
assert result is True
assert isinstance(new_version, DataSet)
assert new_version.name == dataset.name

def test_dataset_validation_invalid_extension(self, unique_name):
"""Test dataset validation with invalid file extension."""
Expand Down
5 changes: 3 additions & 2 deletions sagemaker-train/tests/integ/ai_registry/test_evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -196,8 +196,9 @@ def test_create_evaluator_version(self, unique_name, sample_prompt_file, cleanup
Evaluator.delete_by_name(name=unique_name)
evaluator = Evaluator.create(name=unique_name, type=REWARD_PROMPT, source=sample_prompt_file, wait=False)
# cleanup_list.append(evaluator)
result = evaluator.create_version(source=sample_prompt_file)
assert result is True
new_version = evaluator.create_version(source=sample_prompt_file)
assert isinstance(new_version, Evaluator)
assert new_version.name == evaluator.name
Evaluator.delete_by_name(name=unique_name)

def test_create_reward_prompt_without_source_fails(self, unique_name):
Expand Down
62 changes: 56 additions & 6 deletions sagemaker-train/tests/unit/ai_registry/test_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -374,13 +374,25 @@ def test_create_version_success(self, mock_air_hub, mock_create):
"HubContentDocument": "{}",
"HubContentSearchKeywords": ["customization_technique:sft", "method:generated"]
}
mock_create.return_value = Mock()

dataset = DataSet("test", "arn", "1.0.0", "s3://bucket/prefix", HubContentStatus.AVAILABLE, "desc", CustomizationTechnique.SFT)
session = Mock()
session.sagemaker_config = {"SchemaVersion": "1.0"}
new_dataset = Mock(spec=DataSet)
new_dataset.arn = "test-arn-v2"
new_dataset.version = "2.0.0"
mock_create.return_value = new_dataset

dataset = DataSet(
"test", "arn", "1.0.0", "s3://bucket/prefix", HubContentStatus.AVAILABLE, "desc",
CustomizationTechnique.SFT, sagemaker_session=session,
)
result = dataset.create_version("s3://bucket/new-data")
assert result is True

assert result is new_dataset
mock_create.assert_called_once()
_, create_kwargs = mock_create.call_args
assert create_kwargs["name"] == "test"
assert create_kwargs["source"] == "s3://bucket/new-data"
assert create_kwargs["sagemaker_session"] is session

@patch('sagemaker.ai_registry.dataset.AIRHub')
def test_create_version_failure(self, mock_air_hub):
Expand All @@ -389,7 +401,45 @@ def test_create_version_failure(self, mock_air_hub):
dataset = DataSet("test", "arn", "1.0.0", "s3://bucket/prefix", HubContentStatus.AVAILABLE, "desc", CustomizationTechnique.SFT)
result = dataset.create_version("s3://bucket/new-data")

assert result is False
assert result is None


@patch('sagemaker.train.defaults.resolve_and_validate_role', return_value="arn:aws:iam::123456789012:role/SageMakerRole")
@patch('boto3.client')
@patch('sagemaker.core.helper.session_helper.Session')
@patch('sagemaker.train.common_utils.finetune_utils._get_current_domain_id')
@patch('sagemaker.ai_registry.dataset.DataSet._validate_dataset_file')
@patch('sagemaker.ai_registry.dataset.DataSet._validate_dataset_format')
@patch('sagemaker.ai_registry.dataset.AIRHub')
def test_create_threads_description_to_import_hub_content(
self, mock_air_hub, mock_validate_format, mock_validate_file, mock_get_domain_id,
mock_session, mock_boto_client, mock_resolve_role
):
"""description passed to DataSet.create is forwarded to AIRHub.import_hub_content."""
mock_get_domain_id.return_value = None
mock_session_instance = Mock()
mock_session_instance.get_caller_identity_arn.return_value = "arn:aws:iam::123456789012:role/SageMakerRole"
mock_session.return_value = mock_session_instance
mock_sts_client = Mock()
mock_sts_client.get_caller_identity.return_value = {"Account": "123456789012"}
mock_boto_client.return_value = mock_sts_client
mock_air_hub.import_hub_content.return_value = {"HubContentArn": "test-arn"}
mock_air_hub.describe_hub_content.return_value = {
"HubContentArn": "test-arn",
"HubContentVersion": "1.0.0",
"CreationTime": "2024-01-01",
"LastModifiedTime": "2024-01-01",
}

DataSet.create(
name="test-dataset",
source="s3://bucket/data.jsonl",
description="my description",
wait=False,
)

mock_air_hub.import_hub_content.assert_called_once()
assert mock_air_hub.import_hub_content.call_args.kwargs["description"] == "my description"


class TestDataSetCreateWithContentMetadata:
Expand Down
Loading
Loading