Skip to content
Closed
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
25 changes: 18 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 @@ -519,15 +522,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 +548,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
20 changes: 16 additions & 4 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 Down Expand Up @@ -52,6 +53,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 @@ -495,21 +499,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
23 changes: 17 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,24 @@ 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()
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 +400,7 @@ 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


class TestDataSetCreateWithContentMetadata:
Expand Down
21 changes: 16 additions & 5 deletions sagemaker-train/tests/unit/ai_registry/test_evaluator.py
Original file line number Diff line number Diff line change
Expand Up @@ -192,13 +192,24 @@ def test_get_versions(self, mock_air_hub):

@patch('sagemaker.ai_registry.evaluator.Evaluator.create')
def test_create_version_success(self, mock_create):
mock_create.return_value = MagicMock()

evaluator = Evaluator("test", "1.0.0", "arn", "AWS/Evaluator", method=EvaluatorMethod.LAMBDA, reference="lambda-arn")
session = MagicMock()
new_evaluator = MagicMock()
new_evaluator.arn = "test-arn-v2"
new_evaluator.version = "2.0.0"
mock_create.return_value = new_evaluator

evaluator = Evaluator(
"test", "1.0.0", "arn", "AWS/Evaluator", method=EvaluatorMethod.LAMBDA,
reference="lambda-arn", sagemaker_session=session,
)
result = evaluator.create_version("arn:aws:lambda:us-west-2:123456789012:function:new")
assert result is True

assert result is new_evaluator
mock_create.assert_called_once()
_, create_kwargs = mock_create.call_args
assert create_kwargs["name"] == "test"
assert create_kwargs["source"] == "arn:aws:lambda:us-west-2:123456789012:function:new"
assert create_kwargs["sagemaker_session"] is session

@patch('sagemaker.ai_registry.evaluator.AIRHub')
def test_create_version_failure(self, mock_air_hub):
Expand Down
Loading