Merged
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
30 changes: 18 additions & 12 deletions sagemaker-serve/src/sagemaker/serve/model_builder_servers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -319,43 +319,43 @@ def _build_for_djl(self) -> Model:
logger.debug(f"Using detected notebook instance type: {nb_instance}")

if isinstance(self.model, str) and not self._is_jumpstart_model_id():
# Configure HuggingFace model for DJL
self.env_vars.update({"HF_MODEL_ID": self.model})
# Configure HuggingFace model for DJL (preserve user-provided HF_MODEL_ID)
self.env_vars.setdefault("HF_MODEL_ID", self.model)
Comment thread
aviruthen marked this conversation as resolved.

# Get model configuration for DJL optimization
self.hf_model_config = _get_model_config_properties_from_hf(
self.model, self.env_vars.get("HUGGING_FACE_HUB_TOKEN")
)

# Apply DJL-specific configurations
default_djl_configurations, _default_max_new_tokens = _get_default_djl_configurations(
self.model, self.hf_model_config, self.schema_builder
)
self.env_vars.update(default_djl_configurations)

# Configure schema builder for text generation
if "parameters" not in self.schema_builder.sample_input:
self.schema_builder.sample_input["parameters"] = {}
self.schema_builder.sample_input["parameters"]["max_new_tokens"] = _default_max_new_tokens
# Set DJL serving defaults

# Set DJL serving defaults (only if not already set by user)
djl_env_vars = {
"OPTION_ENGINE": "Python",
"SERVING_MIN_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"OPTION_MODEL_LOADING_TIMEOUT": "240",
"OPTION_PREDICT_TIMEOUT": "60",
"TENSOR_PARALLEL_DEGREE": "1" # Default, will be overridden below
"TENSOR_PARALLEL_DEGREE": "1",
}

# Add HuggingFace authentication
if self.env_vars.get("HUGGING_FACE_HUB_TOKEN"):
djl_env_vars["HF_TOKEN"] = self.env_vars.get("HUGGING_FACE_HUB_TOKEN")

# Update with defaults only if not already set
for key, value in djl_env_vars.items():
self.env_vars.setdefault(key, value)

# DJL downloads models directly from HuggingFace Hub
self.s3_upload_path = None

Expand All@@ -367,6 +367,12 @@ def _build_for_djl(self) -> Model:
else:
self.s3_model_data_url, _ = self._prepare_for_mode()

# Set HF cache env vars to writable location (unconditionally, using setdefault
# to preserve user-provided values). This is needed because /opt/ml/model/ may be
# read-only when source_code artifacts are mounted there.
self.env_vars.setdefault("HF_HOME", "/tmp")
self.env_vars.setdefault("HUGGINGFACE_HUB_CACHE", "/tmp")

# Cache management based on mode
if self.mode in LOCAL_MODES:
self.env_vars.update({"HF_HUB_OFFLINE": "1"})
Expand Down
Empty file.
151 changes: 151 additions & 0 deletions sagemaker-serve/tests/unit/servers/test_djl_hf_cache_env.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,151 @@
"""Tests for DJL builder HF cache environment variables and HF_MODEL_ID handling.
Comment thread
aviruthen marked this conversation as resolved.

Verifies that _build_for_djl() correctly:
- Sets HF_HOME and HUGGINGFACE_HUB_CACHE to /tmp for writable cache
- Preserves user-provided HF_MODEL_ID values (uses setdefault)
- Sets HF_MODEL_ID from model param when not provided by user
- Preserves user-provided HF_HOME and HUGGINGFACE_HUB_CACHE values
"""

import pytest
from unittest.mock import Mock, patch

from sagemaker.serve.model_builder import ModelBuilder
from sagemaker.serve.utils.types import ModelServer
from sagemaker.serve.mode.function_pointers import Mode
from sagemaker.core.resources import Model


MOCK_ROLE_ARN = "arn:aws:iam::000000000000:role/SageMakerRole"
MOCK_IMAGE_URI = "000000000000.dkr.ecr.us-east-1.amazonaws.com/djl-inference:latest"
MOCK_HF_MODEL_CONFIG = {"model_type": "gpt2", "architectures": ["GPT2LMHeadModel"]}


# Common patches needed for _build_for_djl
_DJL_PATCHES = [
"sagemaker.serve.model_builder_servers._get_nb_instance",
"sagemaker.serve.model_builder_servers._get_default_djl_configurations",
"sagemaker.serve.model_builder_servers._get_model_config_properties_from_hf",
"sagemaker.serve.model_builder.ModelBuilder._is_jumpstart_model_id",
"sagemaker.serve.model_builder.ModelBuilder._validate_djl_serving_sample_data",
"sagemaker.serve.model_builder.ModelBuilder._auto_detect_image_uri",
"sagemaker.serve.model_builder.ModelBuilder._prepare_for_mode",
"sagemaker.serve.model_builder.ModelBuilder._create_model",
"sagemaker.serve.model_builder_servers._get_default_tensor_parallel_degree",
"sagemaker.serve.model_builder_servers._get_gpu_info",
]


def _mock_sagemaker_session():
"""Create a mock SageMaker session."""
session = Mock()
session.boto_region_name = "us-east-1"
session.sagemaker_config = {}
session.default_bucket.return_value = "mock-bucket"
session.upload_data.return_value = "s3://mock-bucket/model.tar.gz"
return session


def _create_djl_builder(tmp_path, env_vars=None, mode=Mode.SAGEMAKER_ENDPOINT):
"""Create a ModelBuilder configured for DJL serving tests."""
builder = ModelBuilder(
model="test-org/test-model",
role_arn=MOCK_ROLE_ARN,
sagemaker_session=_mock_sagemaker_session(),
model_path=str(tmp_path),
mode=mode,
image_uri=MOCK_IMAGE_URI,
model_server=ModelServer.DJL_SERVING,
instance_type="ml.g6e.12xlarge",
env_vars=env_vars or {},
)
builder.schema_builder = Mock()
builder.schema_builder.sample_input = {"inputs": "Hello"}
builder._optimizing = False
builder.hf_model_config = MOCK_HF_MODEL_CONFIG
return builder


def _setup_mocks(mocks):
"""Configure common mock return values for DJL build."""
# mocks are in reverse order of _DJL_PATCHES
mock_gpu_info = mocks[-1]
mock_tp_degree = mocks[-2]
mock_create = mocks[-3]
mock_prepare = mocks[-4]
# mock_auto_detect = mocks[-5] # no setup needed
# mock_validate = mocks[-6] # no setup needed
mock_is_js = mocks[-7]
mock_hf_config = mocks[-8]
mock_djl_config = mocks[-9]
mock_nb = mocks[-10]

mock_nb.return_value = None
mock_djl_config.return_value = ({}, 256)
mock_hf_config.return_value = MOCK_HF_MODEL_CONFIG
mock_is_js.return_value = False
mock_prepare.return_value = ("s3://bucket/model", None)
mock_create.return_value = Mock(spec=Model)
mock_tp_degree.return_value = 4
mock_gpu_info.return_value = 4


class TestDjlHfCacheAndModelId:
"""Tests for DJL builder HF cache env vars and HF_MODEL_ID handling."""

@pytest.fixture(autouse=True)
def _patch_djl(self):
"""Apply all DJL-related patches for each test."""
patchers = [patch(p) for p in _DJL_PATCHES]
self._mocks = [p.start() for p in patchers]
_setup_mocks(self._mocks)
yield
for p in patchers:
p.stop()

def test_sets_hf_cache_env_vars_to_tmp(self, tmp_path):
"""HF_HOME and HUGGINGFACE_HUB_CACHE should be /tmp in endpoint mode."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/tmp"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/tmp"

def test_preserves_user_provided_hf_model_id(self, tmp_path):
"""User-provided HF_MODEL_ID must NOT be overridden by model param."""
builder = _create_djl_builder(
tmp_path, env_vars={"HF_MODEL_ID": "/opt/ml/model"}
)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "/opt/ml/model"

def test_sets_hf_model_id_from_model_param_when_not_provided(self, tmp_path):
"""When no user-provided HF_MODEL_ID, it should come from model param."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "test-org/test-model"

def test_preserves_user_provided_hf_cache_dirs(self, tmp_path):
"""User-provided HF_HOME and HUGGINGFACE_HUB_CACHE should be preserved."""
builder = _create_djl_builder(
tmp_path,
env_vars={
"HF_HOME": "/my/custom/cache",
"HUGGINGFACE_HUB_CACHE": "/my/custom/hub",
},
)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/my/custom/cache"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/my/custom/hub"

def test_local_mode_sets_hf_hub_offline(self, tmp_path):
"""HF_HUB_OFFLINE=1 should be set in LOCAL_CONTAINER mode."""
builder = _create_djl_builder(tmp_path, mode=Mode.LOCAL_CONTAINER)
# Local mode doesn't need GPU info mocks for instance_type validation
builder.instance_type = None
builder._build_for_djl()

assert builder.env_vars["HF_HUB_OFFLINE"] == "1"
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Skip to content
Merged
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
30 changes: 18 additions & 12 deletions sagemaker-serve/src/sagemaker/serve/model_builder_servers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -319,43 +319,43 @@ def _build_for_djl(self) -> Model:
logger.debug(f"Using detected notebook instance type: {nb_instance}")

if isinstance(self.model, str) and not self._is_jumpstart_model_id():
# Configure HuggingFace model for DJL
self.env_vars.update({"HF_MODEL_ID": self.model})
# Configure HuggingFace model for DJL (preserve user-provided HF_MODEL_ID)
self.env_vars.setdefault("HF_MODEL_ID", self.model)
Comment thread
aviruthen marked this conversation as resolved.

# Get model configuration for DJL optimization
self.hf_model_config = _get_model_config_properties_from_hf(
self.model, self.env_vars.get("HUGGING_FACE_HUB_TOKEN")
)

# Apply DJL-specific configurations
default_djl_configurations, _default_max_new_tokens = _get_default_djl_configurations(
self.model, self.hf_model_config, self.schema_builder
)
self.env_vars.update(default_djl_configurations)

# Configure schema builder for text generation
if "parameters" not in self.schema_builder.sample_input:
self.schema_builder.sample_input["parameters"] = {}
self.schema_builder.sample_input["parameters"]["max_new_tokens"] = _default_max_new_tokens
# Set DJL serving defaults

# Set DJL serving defaults (only if not already set by user)
djl_env_vars = {
"OPTION_ENGINE": "Python",
"SERVING_MIN_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"OPTION_MODEL_LOADING_TIMEOUT": "240",
"OPTION_PREDICT_TIMEOUT": "60",
"TENSOR_PARALLEL_DEGREE": "1" # Default, will be overridden below
"TENSOR_PARALLEL_DEGREE": "1",
}

# Add HuggingFace authentication
if self.env_vars.get("HUGGING_FACE_HUB_TOKEN"):
djl_env_vars["HF_TOKEN"] = self.env_vars.get("HUGGING_FACE_HUB_TOKEN")

# Update with defaults only if not already set
for key, value in djl_env_vars.items():
self.env_vars.setdefault(key, value)

# DJL downloads models directly from HuggingFace Hub
self.s3_upload_path = None

Expand All@@ -367,6 +367,12 @@ def _build_for_djl(self) -> Model:
else:
self.s3_model_data_url, _ = self._prepare_for_mode()

# Set HF cache env vars to writable location (unconditionally, using setdefault
# to preserve user-provided values). This is needed because /opt/ml/model/ may be
# read-only when source_code artifacts are mounted there.
self.env_vars.setdefault("HF_HOME", "/tmp")
self.env_vars.setdefault("HUGGINGFACE_HUB_CACHE", "/tmp")

# Cache management based on mode
if self.mode in LOCAL_MODES:
self.env_vars.update({"HF_HUB_OFFLINE": "1"})
Expand Down
Empty file.
151 changes: 151 additions & 0 deletions sagemaker-serve/tests/unit/servers/test_djl_hf_cache_env.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,151 @@
"""Tests for DJL builder HF cache environment variables and HF_MODEL_ID handling.
Comment thread
aviruthen marked this conversation as resolved.

Verifies that _build_for_djl() correctly:
- Sets HF_HOME and HUGGINGFACE_HUB_CACHE to /tmp for writable cache
- Preserves user-provided HF_MODEL_ID values (uses setdefault)
- Sets HF_MODEL_ID from model param when not provided by user
- Preserves user-provided HF_HOME and HUGGINGFACE_HUB_CACHE values
"""

import pytest
from unittest.mock import Mock, patch

from sagemaker.serve.model_builder import ModelBuilder
from sagemaker.serve.utils.types import ModelServer
from sagemaker.serve.mode.function_pointers import Mode
from sagemaker.core.resources import Model


MOCK_ROLE_ARN = "arn:aws:iam::000000000000:role/SageMakerRole"
MOCK_IMAGE_URI = "000000000000.dkr.ecr.us-east-1.amazonaws.com/djl-inference:latest"
MOCK_HF_MODEL_CONFIG = {"model_type": "gpt2", "architectures": ["GPT2LMHeadModel"]}


# Common patches needed for _build_for_djl
_DJL_PATCHES = [
"sagemaker.serve.model_builder_servers._get_nb_instance",
"sagemaker.serve.model_builder_servers._get_default_djl_configurations",
"sagemaker.serve.model_builder_servers._get_model_config_properties_from_hf",
"sagemaker.serve.model_builder.ModelBuilder._is_jumpstart_model_id",
"sagemaker.serve.model_builder.ModelBuilder._validate_djl_serving_sample_data",
"sagemaker.serve.model_builder.ModelBuilder._auto_detect_image_uri",
"sagemaker.serve.model_builder.ModelBuilder._prepare_for_mode",
"sagemaker.serve.model_builder.ModelBuilder._create_model",
"sagemaker.serve.model_builder_servers._get_default_tensor_parallel_degree",
"sagemaker.serve.model_builder_servers._get_gpu_info",
]


def _mock_sagemaker_session():
"""Create a mock SageMaker session."""
session = Mock()
session.boto_region_name = "us-east-1"
session.sagemaker_config = {}
session.default_bucket.return_value = "mock-bucket"
session.upload_data.return_value = "s3://mock-bucket/model.tar.gz"
return session


def _create_djl_builder(tmp_path, env_vars=None, mode=Mode.SAGEMAKER_ENDPOINT):
"""Create a ModelBuilder configured for DJL serving tests."""
builder = ModelBuilder(
model="test-org/test-model",
role_arn=MOCK_ROLE_ARN,
sagemaker_session=_mock_sagemaker_session(),
model_path=str(tmp_path),
mode=mode,
image_uri=MOCK_IMAGE_URI,
model_server=ModelServer.DJL_SERVING,
instance_type="ml.g6e.12xlarge",
env_vars=env_vars or {},
)
builder.schema_builder = Mock()
builder.schema_builder.sample_input = {"inputs": "Hello"}
builder._optimizing = False
builder.hf_model_config = MOCK_HF_MODEL_CONFIG
return builder


def _setup_mocks(mocks):
"""Configure common mock return values for DJL build."""
# mocks are in reverse order of _DJL_PATCHES
mock_gpu_info = mocks[-1]
mock_tp_degree = mocks[-2]
mock_create = mocks[-3]
mock_prepare = mocks[-4]
# mock_auto_detect = mocks[-5] # no setup needed
# mock_validate = mocks[-6] # no setup needed
mock_is_js = mocks[-7]
mock_hf_config = mocks[-8]
mock_djl_config = mocks[-9]
mock_nb = mocks[-10]

mock_nb.return_value = None
mock_djl_config.return_value = ({}, 256)
mock_hf_config.return_value = MOCK_HF_MODEL_CONFIG
mock_is_js.return_value = False
mock_prepare.return_value = ("s3://bucket/model", None)
mock_create.return_value = Mock(spec=Model)
mock_tp_degree.return_value = 4
mock_gpu_info.return_value = 4


class TestDjlHfCacheAndModelId:
"""Tests for DJL builder HF cache env vars and HF_MODEL_ID handling."""

@pytest.fixture(autouse=True)
def _patch_djl(self):
"""Apply all DJL-related patches for each test."""
patchers = [patch(p) for p in _DJL_PATCHES]
self._mocks = [p.start() for p in patchers]
_setup_mocks(self._mocks)
yield
for p in patchers:
p.stop()

def test_sets_hf_cache_env_vars_to_tmp(self, tmp_path):
"""HF_HOME and HUGGINGFACE_HUB_CACHE should be /tmp in endpoint mode."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/tmp"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/tmp"

def test_preserves_user_provided_hf_model_id(self, tmp_path):
"""User-provided HF_MODEL_ID must NOT be overridden by model param."""
builder = _create_djl_builder(
tmp_path, env_vars={"HF_MODEL_ID": "/opt/ml/model"}
)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "/opt/ml/model"

def test_sets_hf_model_id_from_model_param_when_not_provided(self, tmp_path):
"""When no user-provided HF_MODEL_ID, it should come from model param."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "test-org/test-model"

def test_preserves_user_provided_hf_cache_dirs(self, tmp_path):
"""User-provided HF_HOME and HUGGINGFACE_HUB_CACHE should be preserved."""
builder = _create_djl_builder(
tmp_path,
env_vars={
"HF_HOME": "/my/custom/cache",
"HUGGINGFACE_HUB_CACHE": "/my/custom/hub",
},
)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/my/custom/cache"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/my/custom/hub"

def test_local_mode_sets_hf_hub_offline(self, tmp_path):
"""HF_HUB_OFFLINE=1 should be set in LOCAL_CONTAINER mode."""
builder = _create_djl_builder(tmp_path, mode=Mode.LOCAL_CONTAINER)
# Local mode doesn't need GPU info mocks for instance_type validation
builder.instance_type = None
builder._build_for_djl()

assert builder.env_vars["HF_HUB_OFFLINE"] == "1"
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
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
30 changes: 18 additions & 12 deletions sagemaker-serve/src/sagemaker/serve/model_builder_servers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -319,43 +319,43 @@ def _build_for_djl(self) -> Model:
logger.debug(f"Using detected notebook instance type: {nb_instance}")

if isinstance(self.model, str) and not self._is_jumpstart_model_id():
# Configure HuggingFace model for DJL
self.env_vars.update({"HF_MODEL_ID": self.model})
# Configure HuggingFace model for DJL (preserve user-provided HF_MODEL_ID)
self.env_vars.setdefault("HF_MODEL_ID", self.model)
Comment thread
aviruthen marked this conversation as resolved.

# Get model configuration for DJL optimization
self.hf_model_config = _get_model_config_properties_from_hf(
self.model, self.env_vars.get("HUGGING_FACE_HUB_TOKEN")
)

# Apply DJL-specific configurations
default_djl_configurations, _default_max_new_tokens = _get_default_djl_configurations(
self.model, self.hf_model_config, self.schema_builder
)
self.env_vars.update(default_djl_configurations)

# Configure schema builder for text generation
if "parameters" not in self.schema_builder.sample_input:
self.schema_builder.sample_input["parameters"] = {}
self.schema_builder.sample_input["parameters"]["max_new_tokens"] = _default_max_new_tokens
# Set DJL serving defaults

# Set DJL serving defaults (only if not already set by user)
djl_env_vars = {
"OPTION_ENGINE": "Python",
"SERVING_MIN_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"OPTION_MODEL_LOADING_TIMEOUT": "240",
"OPTION_PREDICT_TIMEOUT": "60",
"TENSOR_PARALLEL_DEGREE": "1" # Default, will be overridden below
"TENSOR_PARALLEL_DEGREE": "1",
}

# Add HuggingFace authentication
if self.env_vars.get("HUGGING_FACE_HUB_TOKEN"):
djl_env_vars["HF_TOKEN"] = self.env_vars.get("HUGGING_FACE_HUB_TOKEN")

# Update with defaults only if not already set
for key, value in djl_env_vars.items():
self.env_vars.setdefault(key, value)

# DJL downloads models directly from HuggingFace Hub
self.s3_upload_path = None

Expand All@@ -367,6 +367,12 @@ def _build_for_djl(self) -> Model:
else:
self.s3_model_data_url, _ = self._prepare_for_mode()

# Set HF cache env vars to writable location (unconditionally, using setdefault
# to preserve user-provided values). This is needed because /opt/ml/model/ may be
# read-only when source_code artifacts are mounted there.
self.env_vars.setdefault("HF_HOME", "/tmp")
self.env_vars.setdefault("HUGGINGFACE_HUB_CACHE", "/tmp")

# Cache management based on mode
if self.mode in LOCAL_MODES:
self.env_vars.update({"HF_HUB_OFFLINE": "1"})
Expand Down
Empty file.
151 changes: 151 additions & 0 deletions sagemaker-serve/tests/unit/servers/test_djl_hf_cache_env.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,151 @@
"""Tests for DJL builder HF cache environment variables and HF_MODEL_ID handling.
Comment thread
aviruthen marked this conversation as resolved.

Verifies that _build_for_djl() correctly:
- Sets HF_HOME and HUGGINGFACE_HUB_CACHE to /tmp for writable cache
- Preserves user-provided HF_MODEL_ID values (uses setdefault)
- Sets HF_MODEL_ID from model param when not provided by user
- Preserves user-provided HF_HOME and HUGGINGFACE_HUB_CACHE values
"""

import pytest
from unittest.mock import Mock, patch

from sagemaker.serve.model_builder import ModelBuilder
from sagemaker.serve.utils.types import ModelServer
from sagemaker.serve.mode.function_pointers import Mode
from sagemaker.core.resources import Model


MOCK_ROLE_ARN = "arn:aws:iam::000000000000:role/SageMakerRole"
MOCK_IMAGE_URI = "000000000000.dkr.ecr.us-east-1.amazonaws.com/djl-inference:latest"
MOCK_HF_MODEL_CONFIG = {"model_type": "gpt2", "architectures": ["GPT2LMHeadModel"]}


# Common patches needed for _build_for_djl
_DJL_PATCHES = [
"sagemaker.serve.model_builder_servers._get_nb_instance",
"sagemaker.serve.model_builder_servers._get_default_djl_configurations",
"sagemaker.serve.model_builder_servers._get_model_config_properties_from_hf",
"sagemaker.serve.model_builder.ModelBuilder._is_jumpstart_model_id",
"sagemaker.serve.model_builder.ModelBuilder._validate_djl_serving_sample_data",
"sagemaker.serve.model_builder.ModelBuilder._auto_detect_image_uri",
"sagemaker.serve.model_builder.ModelBuilder._prepare_for_mode",
"sagemaker.serve.model_builder.ModelBuilder._create_model",
"sagemaker.serve.model_builder_servers._get_default_tensor_parallel_degree",
"sagemaker.serve.model_builder_servers._get_gpu_info",
]


def _mock_sagemaker_session():
"""Create a mock SageMaker session."""
session = Mock()
session.boto_region_name = "us-east-1"
session.sagemaker_config = {}
session.default_bucket.return_value = "mock-bucket"
session.upload_data.return_value = "s3://mock-bucket/model.tar.gz"
return session


def _create_djl_builder(tmp_path, env_vars=None, mode=Mode.SAGEMAKER_ENDPOINT):
"""Create a ModelBuilder configured for DJL serving tests."""
builder = ModelBuilder(
model="test-org/test-model",
role_arn=MOCK_ROLE_ARN,
sagemaker_session=_mock_sagemaker_session(),
model_path=str(tmp_path),
mode=mode,
image_uri=MOCK_IMAGE_URI,
model_server=ModelServer.DJL_SERVING,
instance_type="ml.g6e.12xlarge",
env_vars=env_vars or {},
)
builder.schema_builder = Mock()
builder.schema_builder.sample_input = {"inputs": "Hello"}
builder._optimizing = False
builder.hf_model_config = MOCK_HF_MODEL_CONFIG
return builder


def _setup_mocks(mocks):
"""Configure common mock return values for DJL build."""
# mocks are in reverse order of _DJL_PATCHES
mock_gpu_info = mocks[-1]
mock_tp_degree = mocks[-2]
mock_create = mocks[-3]
mock_prepare = mocks[-4]
# mock_auto_detect = mocks[-5] # no setup needed
# mock_validate = mocks[-6] # no setup needed
mock_is_js = mocks[-7]
mock_hf_config = mocks[-8]
mock_djl_config = mocks[-9]
mock_nb = mocks[-10]

mock_nb.return_value = None
mock_djl_config.return_value = ({}, 256)
mock_hf_config.return_value = MOCK_HF_MODEL_CONFIG
mock_is_js.return_value = False
mock_prepare.return_value = ("s3://bucket/model", None)
mock_create.return_value = Mock(spec=Model)
mock_tp_degree.return_value = 4
mock_gpu_info.return_value = 4


class TestDjlHfCacheAndModelId:
"""Tests for DJL builder HF cache env vars and HF_MODEL_ID handling."""

@pytest.fixture(autouse=True)
def _patch_djl(self):
"""Apply all DJL-related patches for each test."""
patchers = [patch(p) for p in _DJL_PATCHES]
self._mocks = [p.start() for p in patchers]
_setup_mocks(self._mocks)
yield
for p in patchers:
p.stop()

def test_sets_hf_cache_env_vars_to_tmp(self, tmp_path):
"""HF_HOME and HUGGINGFACE_HUB_CACHE should be /tmp in endpoint mode."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/tmp"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/tmp"

def test_preserves_user_provided_hf_model_id(self, tmp_path):
"""User-provided HF_MODEL_ID must NOT be overridden by model param."""
builder = _create_djl_builder(
tmp_path, env_vars={"HF_MODEL_ID": "/opt/ml/model"}
)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "/opt/ml/model"

def test_sets_hf_model_id_from_model_param_when_not_provided(self, tmp_path):
"""When no user-provided HF_MODEL_ID, it should come from model param."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "test-org/test-model"

def test_preserves_user_provided_hf_cache_dirs(self, tmp_path):
"""User-provided HF_HOME and HUGGINGFACE_HUB_CACHE should be preserved."""
builder = _create_djl_builder(
tmp_path,
env_vars={
"HF_HOME": "/my/custom/cache",
"HUGGINGFACE_HUB_CACHE": "/my/custom/hub",
},
)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/my/custom/cache"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/my/custom/hub"

def test_local_mode_sets_hf_hub_offline(self, tmp_path):
"""HF_HUB_OFFLINE=1 should be set in LOCAL_CONTAINER mode."""
builder = _create_djl_builder(tmp_path, mode=Mode.LOCAL_CONTAINER)
# Local mode doesn't need GPU info mocks for instance_type validation
builder.instance_type = None
builder._build_for_djl()

assert builder.env_vars["HF_HUB_OFFLINE"] == "1"
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
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
30 changes: 18 additions & 12 deletions sagemaker-serve/src/sagemaker/serve/model_builder_servers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -319,43 +319,43 @@ def _build_for_djl(self) -> Model:
logger.debug(f"Using detected notebook instance type: {nb_instance}")

if isinstance(self.model, str) and not self._is_jumpstart_model_id():
# Configure HuggingFace model for DJL
self.env_vars.update({"HF_MODEL_ID": self.model})
# Configure HuggingFace model for DJL (preserve user-provided HF_MODEL_ID)
self.env_vars.setdefault("HF_MODEL_ID", self.model)
Comment thread
aviruthen marked this conversation as resolved.

# Get model configuration for DJL optimization
self.hf_model_config = _get_model_config_properties_from_hf(
self.model, self.env_vars.get("HUGGING_FACE_HUB_TOKEN")
)

# Apply DJL-specific configurations
default_djl_configurations, _default_max_new_tokens = _get_default_djl_configurations(
self.model, self.hf_model_config, self.schema_builder
)
self.env_vars.update(default_djl_configurations)

# Configure schema builder for text generation
if "parameters" not in self.schema_builder.sample_input:
self.schema_builder.sample_input["parameters"] = {}
self.schema_builder.sample_input["parameters"]["max_new_tokens"] = _default_max_new_tokens
# Set DJL serving defaults

# Set DJL serving defaults (only if not already set by user)
djl_env_vars = {
"OPTION_ENGINE": "Python",
"SERVING_MIN_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"OPTION_MODEL_LOADING_TIMEOUT": "240",
"OPTION_PREDICT_TIMEOUT": "60",
"TENSOR_PARALLEL_DEGREE": "1" # Default, will be overridden below
"TENSOR_PARALLEL_DEGREE": "1",
}

# Add HuggingFace authentication
if self.env_vars.get("HUGGING_FACE_HUB_TOKEN"):
djl_env_vars["HF_TOKEN"] = self.env_vars.get("HUGGING_FACE_HUB_TOKEN")

# Update with defaults only if not already set
for key, value in djl_env_vars.items():
self.env_vars.setdefault(key, value)

# DJL downloads models directly from HuggingFace Hub
self.s3_upload_path = None

Expand All@@ -367,6 +367,12 @@ def _build_for_djl(self) -> Model:
else:
self.s3_model_data_url, _ = self._prepare_for_mode()

# Set HF cache env vars to writable location (unconditionally, using setdefault
# to preserve user-provided values). This is needed because /opt/ml/model/ may be
# read-only when source_code artifacts are mounted there.
self.env_vars.setdefault("HF_HOME", "/tmp")
self.env_vars.setdefault("HUGGINGFACE_HUB_CACHE", "/tmp")

# Cache management based on mode
if self.mode in LOCAL_MODES:
self.env_vars.update({"HF_HUB_OFFLINE": "1"})
Expand Down
Empty file.
151 changes: 151 additions & 0 deletions sagemaker-serve/tests/unit/servers/test_djl_hf_cache_env.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,151 @@
"""Tests for DJL builder HF cache environment variables and HF_MODEL_ID handling.
Comment thread
aviruthen marked this conversation as resolved.

Verifies that _build_for_djl() correctly:
- Sets HF_HOME and HUGGINGFACE_HUB_CACHE to /tmp for writable cache
- Preserves user-provided HF_MODEL_ID values (uses setdefault)
- Sets HF_MODEL_ID from model param when not provided by user
- Preserves user-provided HF_HOME and HUGGINGFACE_HUB_CACHE values
"""

import pytest
from unittest.mock import Mock, patch

from sagemaker.serve.model_builder import ModelBuilder
from sagemaker.serve.utils.types import ModelServer
from sagemaker.serve.mode.function_pointers import Mode
from sagemaker.core.resources import Model


MOCK_ROLE_ARN = "arn:aws:iam::000000000000:role/SageMakerRole"
MOCK_IMAGE_URI = "000000000000.dkr.ecr.us-east-1.amazonaws.com/djl-inference:latest"
MOCK_HF_MODEL_CONFIG = {"model_type": "gpt2", "architectures": ["GPT2LMHeadModel"]}


# Common patches needed for _build_for_djl
_DJL_PATCHES = [
"sagemaker.serve.model_builder_servers._get_nb_instance",
"sagemaker.serve.model_builder_servers._get_default_djl_configurations",
"sagemaker.serve.model_builder_servers._get_model_config_properties_from_hf",
"sagemaker.serve.model_builder.ModelBuilder._is_jumpstart_model_id",
"sagemaker.serve.model_builder.ModelBuilder._validate_djl_serving_sample_data",
"sagemaker.serve.model_builder.ModelBuilder._auto_detect_image_uri",
"sagemaker.serve.model_builder.ModelBuilder._prepare_for_mode",
"sagemaker.serve.model_builder.ModelBuilder._create_model",
"sagemaker.serve.model_builder_servers._get_default_tensor_parallel_degree",
"sagemaker.serve.model_builder_servers._get_gpu_info",
]


def _mock_sagemaker_session():
"""Create a mock SageMaker session."""
session = Mock()
session.boto_region_name = "us-east-1"
session.sagemaker_config = {}
session.default_bucket.return_value = "mock-bucket"
session.upload_data.return_value = "s3://mock-bucket/model.tar.gz"
return session


def _create_djl_builder(tmp_path, env_vars=None, mode=Mode.SAGEMAKER_ENDPOINT):
"""Create a ModelBuilder configured for DJL serving tests."""
builder = ModelBuilder(
model="test-org/test-model",
role_arn=MOCK_ROLE_ARN,
sagemaker_session=_mock_sagemaker_session(),
model_path=str(tmp_path),
mode=mode,
image_uri=MOCK_IMAGE_URI,
model_server=ModelServer.DJL_SERVING,
instance_type="ml.g6e.12xlarge",
env_vars=env_vars or {},
)
builder.schema_builder = Mock()
builder.schema_builder.sample_input = {"inputs": "Hello"}
builder._optimizing = False
builder.hf_model_config = MOCK_HF_MODEL_CONFIG
return builder


def _setup_mocks(mocks):
"""Configure common mock return values for DJL build."""
# mocks are in reverse order of _DJL_PATCHES
mock_gpu_info = mocks[-1]
mock_tp_degree = mocks[-2]
mock_create = mocks[-3]
mock_prepare = mocks[-4]
# mock_auto_detect = mocks[-5] # no setup needed
# mock_validate = mocks[-6] # no setup needed
mock_is_js = mocks[-7]
mock_hf_config = mocks[-8]
mock_djl_config = mocks[-9]
mock_nb = mocks[-10]

mock_nb.return_value = None
mock_djl_config.return_value = ({}, 256)
mock_hf_config.return_value = MOCK_HF_MODEL_CONFIG
mock_is_js.return_value = False
mock_prepare.return_value = ("s3://bucket/model", None)
mock_create.return_value = Mock(spec=Model)
mock_tp_degree.return_value = 4
mock_gpu_info.return_value = 4


class TestDjlHfCacheAndModelId:
"""Tests for DJL builder HF cache env vars and HF_MODEL_ID handling."""

@pytest.fixture(autouse=True)
def _patch_djl(self):
"""Apply all DJL-related patches for each test."""
patchers = [patch(p) for p in _DJL_PATCHES]
self._mocks = [p.start() for p in patchers]
_setup_mocks(self._mocks)
yield
for p in patchers:
p.stop()

def test_sets_hf_cache_env_vars_to_tmp(self, tmp_path):
"""HF_HOME and HUGGINGFACE_HUB_CACHE should be /tmp in endpoint mode."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/tmp"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/tmp"

def test_preserves_user_provided_hf_model_id(self, tmp_path):
"""User-provided HF_MODEL_ID must NOT be overridden by model param."""
builder = _create_djl_builder(
tmp_path, env_vars={"HF_MODEL_ID": "/opt/ml/model"}
)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "/opt/ml/model"

def test_sets_hf_model_id_from_model_param_when_not_provided(self, tmp_path):
"""When no user-provided HF_MODEL_ID, it should come from model param."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "test-org/test-model"

def test_preserves_user_provided_hf_cache_dirs(self, tmp_path):
"""User-provided HF_HOME and HUGGINGFACE_HUB_CACHE should be preserved."""
builder = _create_djl_builder(
tmp_path,
env_vars={
"HF_HOME": "/my/custom/cache",
"HUGGINGFACE_HUB_CACHE": "/my/custom/hub",
},
)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/my/custom/cache"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/my/custom/hub"

def test_local_mode_sets_hf_hub_offline(self, tmp_path):
"""HF_HUB_OFFLINE=1 should be set in LOCAL_CONTAINER mode."""
builder = _create_djl_builder(tmp_path, mode=Mode.LOCAL_CONTAINER)
# Local mode doesn't need GPU info mocks for instance_type validation
builder.instance_type = None
builder._build_for_djl()

assert builder.env_vars["HF_HUB_OFFLINE"] == "1"
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
Skip to content
Merged
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
30 changes: 18 additions & 12 deletions sagemaker-serve/src/sagemaker/serve/model_builder_servers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -319,43 +319,43 @@ def _build_for_djl(self) -> Model:
logger.debug(f"Using detected notebook instance type: {nb_instance}")

if isinstance(self.model, str) and not self._is_jumpstart_model_id():
# Configure HuggingFace model for DJL
self.env_vars.update({"HF_MODEL_ID": self.model})
# Configure HuggingFace model for DJL (preserve user-provided HF_MODEL_ID)
self.env_vars.setdefault("HF_MODEL_ID", self.model)
Comment thread
aviruthen marked this conversation as resolved.

# Get model configuration for DJL optimization
self.hf_model_config = _get_model_config_properties_from_hf(
self.model, self.env_vars.get("HUGGING_FACE_HUB_TOKEN")
)

# Apply DJL-specific configurations
default_djl_configurations, _default_max_new_tokens = _get_default_djl_configurations(
self.model, self.hf_model_config, self.schema_builder
)
self.env_vars.update(default_djl_configurations)

# Configure schema builder for text generation
if "parameters" not in self.schema_builder.sample_input:
self.schema_builder.sample_input["parameters"] = {}
self.schema_builder.sample_input["parameters"]["max_new_tokens"] = _default_max_new_tokens
# Set DJL serving defaults

# Set DJL serving defaults (only if not already set by user)
djl_env_vars = {
"OPTION_ENGINE": "Python",
"SERVING_MIN_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"OPTION_MODEL_LOADING_TIMEOUT": "240",
"OPTION_PREDICT_TIMEOUT": "60",
"TENSOR_PARALLEL_DEGREE": "1" # Default, will be overridden below
"TENSOR_PARALLEL_DEGREE": "1",
}

# Add HuggingFace authentication
if self.env_vars.get("HUGGING_FACE_HUB_TOKEN"):
djl_env_vars["HF_TOKEN"] = self.env_vars.get("HUGGING_FACE_HUB_TOKEN")

# Update with defaults only if not already set
for key, value in djl_env_vars.items():
self.env_vars.setdefault(key, value)

# DJL downloads models directly from HuggingFace Hub
self.s3_upload_path = None

Expand All@@ -367,6 +367,12 @@ def _build_for_djl(self) -> Model:
else:
self.s3_model_data_url, _ = self._prepare_for_mode()

# Set HF cache env vars to writable location (unconditionally, using setdefault
# to preserve user-provided values). This is needed because /opt/ml/model/ may be
# read-only when source_code artifacts are mounted there.
self.env_vars.setdefault("HF_HOME", "/tmp")
self.env_vars.setdefault("HUGGINGFACE_HUB_CACHE", "/tmp")

# Cache management based on mode
if self.mode in LOCAL_MODES:
self.env_vars.update({"HF_HUB_OFFLINE": "1"})
Expand Down
Empty file.
151 changes: 151 additions & 0 deletions sagemaker-serve/tests/unit/servers/test_djl_hf_cache_env.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,151 @@
"""Tests for DJL builder HF cache environment variables and HF_MODEL_ID handling.
Comment thread
aviruthen marked this conversation as resolved.

Verifies that _build_for_djl() correctly:
- Sets HF_HOME and HUGGINGFACE_HUB_CACHE to /tmp for writable cache
- Preserves user-provided HF_MODEL_ID values (uses setdefault)
- Sets HF_MODEL_ID from model param when not provided by user
- Preserves user-provided HF_HOME and HUGGINGFACE_HUB_CACHE values
"""

import pytest
from unittest.mock import Mock, patch

from sagemaker.serve.model_builder import ModelBuilder
from sagemaker.serve.utils.types import ModelServer
from sagemaker.serve.mode.function_pointers import Mode
from sagemaker.core.resources import Model


MOCK_ROLE_ARN = "arn:aws:iam::000000000000:role/SageMakerRole"
MOCK_IMAGE_URI = "000000000000.dkr.ecr.us-east-1.amazonaws.com/djl-inference:latest"
MOCK_HF_MODEL_CONFIG = {"model_type": "gpt2", "architectures": ["GPT2LMHeadModel"]}


# Common patches needed for _build_for_djl
_DJL_PATCHES = [
"sagemaker.serve.model_builder_servers._get_nb_instance",
"sagemaker.serve.model_builder_servers._get_default_djl_configurations",
"sagemaker.serve.model_builder_servers._get_model_config_properties_from_hf",
"sagemaker.serve.model_builder.ModelBuilder._is_jumpstart_model_id",
"sagemaker.serve.model_builder.ModelBuilder._validate_djl_serving_sample_data",
"sagemaker.serve.model_builder.ModelBuilder._auto_detect_image_uri",
"sagemaker.serve.model_builder.ModelBuilder._prepare_for_mode",
"sagemaker.serve.model_builder.ModelBuilder._create_model",
"sagemaker.serve.model_builder_servers._get_default_tensor_parallel_degree",
"sagemaker.serve.model_builder_servers._get_gpu_info",
]


def _mock_sagemaker_session():
"""Create a mock SageMaker session."""
session = Mock()
session.boto_region_name = "us-east-1"
session.sagemaker_config = {}
session.default_bucket.return_value = "mock-bucket"
session.upload_data.return_value = "s3://mock-bucket/model.tar.gz"
return session


def _create_djl_builder(tmp_path, env_vars=None, mode=Mode.SAGEMAKER_ENDPOINT):
"""Create a ModelBuilder configured for DJL serving tests."""
builder = ModelBuilder(
model="test-org/test-model",
role_arn=MOCK_ROLE_ARN,
sagemaker_session=_mock_sagemaker_session(),
model_path=str(tmp_path),
mode=mode,
image_uri=MOCK_IMAGE_URI,
model_server=ModelServer.DJL_SERVING,
instance_type="ml.g6e.12xlarge",
env_vars=env_vars or {},
)
builder.schema_builder = Mock()
builder.schema_builder.sample_input = {"inputs": "Hello"}
builder._optimizing = False
builder.hf_model_config = MOCK_HF_MODEL_CONFIG
return builder


def _setup_mocks(mocks):
"""Configure common mock return values for DJL build."""
# mocks are in reverse order of _DJL_PATCHES
mock_gpu_info = mocks[-1]
mock_tp_degree = mocks[-2]
mock_create = mocks[-3]
mock_prepare = mocks[-4]
# mock_auto_detect = mocks[-5] # no setup needed
# mock_validate = mocks[-6] # no setup needed
mock_is_js = mocks[-7]
mock_hf_config = mocks[-8]
mock_djl_config = mocks[-9]
mock_nb = mocks[-10]

mock_nb.return_value = None
mock_djl_config.return_value = ({}, 256)
mock_hf_config.return_value = MOCK_HF_MODEL_CONFIG
mock_is_js.return_value = False
mock_prepare.return_value = ("s3://bucket/model", None)
mock_create.return_value = Mock(spec=Model)
mock_tp_degree.return_value = 4
mock_gpu_info.return_value = 4


class TestDjlHfCacheAndModelId:
"""Tests for DJL builder HF cache env vars and HF_MODEL_ID handling."""

@pytest.fixture(autouse=True)
def _patch_djl(self):
"""Apply all DJL-related patches for each test."""
patchers = [patch(p) for p in _DJL_PATCHES]
self._mocks = [p.start() for p in patchers]
_setup_mocks(self._mocks)
yield
for p in patchers:
p.stop()

def test_sets_hf_cache_env_vars_to_tmp(self, tmp_path):
"""HF_HOME and HUGGINGFACE_HUB_CACHE should be /tmp in endpoint mode."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/tmp"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/tmp"

def test_preserves_user_provided_hf_model_id(self, tmp_path):
"""User-provided HF_MODEL_ID must NOT be overridden by model param."""
builder = _create_djl_builder(
tmp_path, env_vars={"HF_MODEL_ID": "/opt/ml/model"}
)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "/opt/ml/model"

def test_sets_hf_model_id_from_model_param_when_not_provided(self, tmp_path):
"""When no user-provided HF_MODEL_ID, it should come from model param."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "test-org/test-model"

def test_preserves_user_provided_hf_cache_dirs(self, tmp_path):
"""User-provided HF_HOME and HUGGINGFACE_HUB_CACHE should be preserved."""
builder = _create_djl_builder(
tmp_path,
env_vars={
"HF_HOME": "/my/custom/cache",
"HUGGINGFACE_HUB_CACHE": "/my/custom/hub",
},
)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/my/custom/cache"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/my/custom/hub"

def test_local_mode_sets_hf_hub_offline(self, tmp_path):
"""HF_HUB_OFFLINE=1 should be set in LOCAL_CONTAINER mode."""
builder = _create_djl_builder(tmp_path, mode=Mode.LOCAL_CONTAINER)
# Local mode doesn't need GPU info mocks for instance_type validation
builder.instance_type = None
builder._build_for_djl()

assert builder.env_vars["HF_HUB_OFFLINE"] == "1"
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
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
30 changes: 18 additions & 12 deletions sagemaker-serve/src/sagemaker/serve/model_builder_servers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -319,43 +319,43 @@ def _build_for_djl(self) -> Model:
logger.debug(f"Using detected notebook instance type: {nb_instance}")

if isinstance(self.model, str) and not self._is_jumpstart_model_id():
# Configure HuggingFace model for DJL
self.env_vars.update({"HF_MODEL_ID": self.model})
# Configure HuggingFace model for DJL (preserve user-provided HF_MODEL_ID)
self.env_vars.setdefault("HF_MODEL_ID", self.model)
Comment thread
aviruthen marked this conversation as resolved.

# Get model configuration for DJL optimization
self.hf_model_config = _get_model_config_properties_from_hf(
self.model, self.env_vars.get("HUGGING_FACE_HUB_TOKEN")
)

# Apply DJL-specific configurations
default_djl_configurations, _default_max_new_tokens = _get_default_djl_configurations(
self.model, self.hf_model_config, self.schema_builder
)
self.env_vars.update(default_djl_configurations)

# Configure schema builder for text generation
if "parameters" not in self.schema_builder.sample_input:
self.schema_builder.sample_input["parameters"] = {}
self.schema_builder.sample_input["parameters"]["max_new_tokens"] = _default_max_new_tokens
# Set DJL serving defaults

# Set DJL serving defaults (only if not already set by user)
djl_env_vars = {
"OPTION_ENGINE": "Python",
"SERVING_MIN_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"OPTION_MODEL_LOADING_TIMEOUT": "240",
"OPTION_PREDICT_TIMEOUT": "60",
"TENSOR_PARALLEL_DEGREE": "1" # Default, will be overridden below
"TENSOR_PARALLEL_DEGREE": "1",
}

# Add HuggingFace authentication
if self.env_vars.get("HUGGING_FACE_HUB_TOKEN"):
djl_env_vars["HF_TOKEN"] = self.env_vars.get("HUGGING_FACE_HUB_TOKEN")

# Update with defaults only if not already set
for key, value in djl_env_vars.items():
self.env_vars.setdefault(key, value)

# DJL downloads models directly from HuggingFace Hub
self.s3_upload_path = None

Expand All@@ -367,6 +367,12 @@ def _build_for_djl(self) -> Model:
else:
self.s3_model_data_url, _ = self._prepare_for_mode()

# Set HF cache env vars to writable location (unconditionally, using setdefault
# to preserve user-provided values). This is needed because /opt/ml/model/ may be
# read-only when source_code artifacts are mounted there.
self.env_vars.setdefault("HF_HOME", "/tmp")
self.env_vars.setdefault("HUGGINGFACE_HUB_CACHE", "/tmp")

# Cache management based on mode
if self.mode in LOCAL_MODES:
self.env_vars.update({"HF_HUB_OFFLINE": "1"})
Expand Down
Empty file.
151 changes: 151 additions & 0 deletions sagemaker-serve/tests/unit/servers/test_djl_hf_cache_env.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,151 @@
"""Tests for DJL builder HF cache environment variables and HF_MODEL_ID handling.
Comment thread
aviruthen marked this conversation as resolved.

Verifies that _build_for_djl() correctly:
- Sets HF_HOME and HUGGINGFACE_HUB_CACHE to /tmp for writable cache
- Preserves user-provided HF_MODEL_ID values (uses setdefault)
- Sets HF_MODEL_ID from model param when not provided by user
- Preserves user-provided HF_HOME and HUGGINGFACE_HUB_CACHE values
"""

import pytest
from unittest.mock import Mock, patch

from sagemaker.serve.model_builder import ModelBuilder
from sagemaker.serve.utils.types import ModelServer
from sagemaker.serve.mode.function_pointers import Mode
from sagemaker.core.resources import Model


MOCK_ROLE_ARN = "arn:aws:iam::000000000000:role/SageMakerRole"
MOCK_IMAGE_URI = "000000000000.dkr.ecr.us-east-1.amazonaws.com/djl-inference:latest"
MOCK_HF_MODEL_CONFIG = {"model_type": "gpt2", "architectures": ["GPT2LMHeadModel"]}


# Common patches needed for _build_for_djl
_DJL_PATCHES = [
"sagemaker.serve.model_builder_servers._get_nb_instance",
"sagemaker.serve.model_builder_servers._get_default_djl_configurations",
"sagemaker.serve.model_builder_servers._get_model_config_properties_from_hf",
"sagemaker.serve.model_builder.ModelBuilder._is_jumpstart_model_id",
"sagemaker.serve.model_builder.ModelBuilder._validate_djl_serving_sample_data",
"sagemaker.serve.model_builder.ModelBuilder._auto_detect_image_uri",
"sagemaker.serve.model_builder.ModelBuilder._prepare_for_mode",
"sagemaker.serve.model_builder.ModelBuilder._create_model",
"sagemaker.serve.model_builder_servers._get_default_tensor_parallel_degree",
"sagemaker.serve.model_builder_servers._get_gpu_info",
]


def _mock_sagemaker_session():
"""Create a mock SageMaker session."""
session = Mock()
session.boto_region_name = "us-east-1"
session.sagemaker_config = {}
session.default_bucket.return_value = "mock-bucket"
session.upload_data.return_value = "s3://mock-bucket/model.tar.gz"
return session


def _create_djl_builder(tmp_path, env_vars=None, mode=Mode.SAGEMAKER_ENDPOINT):
"""Create a ModelBuilder configured for DJL serving tests."""
builder = ModelBuilder(
model="test-org/test-model",
role_arn=MOCK_ROLE_ARN,
sagemaker_session=_mock_sagemaker_session(),
model_path=str(tmp_path),
mode=mode,
image_uri=MOCK_IMAGE_URI,
model_server=ModelServer.DJL_SERVING,
instance_type="ml.g6e.12xlarge",
env_vars=env_vars or {},
)
builder.schema_builder = Mock()
builder.schema_builder.sample_input = {"inputs": "Hello"}
builder._optimizing = False
builder.hf_model_config = MOCK_HF_MODEL_CONFIG
return builder


def _setup_mocks(mocks):
"""Configure common mock return values for DJL build."""
# mocks are in reverse order of _DJL_PATCHES
mock_gpu_info = mocks[-1]
mock_tp_degree = mocks[-2]
mock_create = mocks[-3]
mock_prepare = mocks[-4]
# mock_auto_detect = mocks[-5] # no setup needed
# mock_validate = mocks[-6] # no setup needed
mock_is_js = mocks[-7]
mock_hf_config = mocks[-8]
mock_djl_config = mocks[-9]
mock_nb = mocks[-10]

mock_nb.return_value = None
mock_djl_config.return_value = ({}, 256)
mock_hf_config.return_value = MOCK_HF_MODEL_CONFIG
mock_is_js.return_value = False
mock_prepare.return_value = ("s3://bucket/model", None)
mock_create.return_value = Mock(spec=Model)
mock_tp_degree.return_value = 4
mock_gpu_info.return_value = 4


class TestDjlHfCacheAndModelId:
"""Tests for DJL builder HF cache env vars and HF_MODEL_ID handling."""

@pytest.fixture(autouse=True)
def _patch_djl(self):
"""Apply all DJL-related patches for each test."""
patchers = [patch(p) for p in _DJL_PATCHES]
self._mocks = [p.start() for p in patchers]
_setup_mocks(self._mocks)
yield
for p in patchers:
p.stop()

def test_sets_hf_cache_env_vars_to_tmp(self, tmp_path):
"""HF_HOME and HUGGINGFACE_HUB_CACHE should be /tmp in endpoint mode."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/tmp"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/tmp"

def test_preserves_user_provided_hf_model_id(self, tmp_path):
"""User-provided HF_MODEL_ID must NOT be overridden by model param."""
builder = _create_djl_builder(
tmp_path, env_vars={"HF_MODEL_ID": "/opt/ml/model"}
)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "/opt/ml/model"

def test_sets_hf_model_id_from_model_param_when_not_provided(self, tmp_path):
"""When no user-provided HF_MODEL_ID, it should come from model param."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "test-org/test-model"

def test_preserves_user_provided_hf_cache_dirs(self, tmp_path):
"""User-provided HF_HOME and HUGGINGFACE_HUB_CACHE should be preserved."""
builder = _create_djl_builder(
tmp_path,
env_vars={
"HF_HOME": "/my/custom/cache",
"HUGGINGFACE_HUB_CACHE": "/my/custom/hub",
},
)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/my/custom/cache"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/my/custom/hub"

def test_local_mode_sets_hf_hub_offline(self, tmp_path):
"""HF_HUB_OFFLINE=1 should be set in LOCAL_CONTAINER mode."""
builder = _create_djl_builder(tmp_path, mode=Mode.LOCAL_CONTAINER)
# Local mode doesn't need GPU info mocks for instance_type validation
builder.instance_type = None
builder._build_for_djl()

assert builder.env_vars["HF_HUB_OFFLINE"] == "1"
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
Merged
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
30 changes: 18 additions & 12 deletions sagemaker-serve/src/sagemaker/serve/model_builder_servers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -319,43 +319,43 @@ def _build_for_djl(self) -> Model:
logger.debug(f"Using detected notebook instance type: {nb_instance}")

if isinstance(self.model, str) and not self._is_jumpstart_model_id():
# Configure HuggingFace model for DJL
self.env_vars.update({"HF_MODEL_ID": self.model})
# Configure HuggingFace model for DJL (preserve user-provided HF_MODEL_ID)
self.env_vars.setdefault("HF_MODEL_ID", self.model)
Comment thread
aviruthen marked this conversation as resolved.

# Get model configuration for DJL optimization
self.hf_model_config = _get_model_config_properties_from_hf(
self.model, self.env_vars.get("HUGGING_FACE_HUB_TOKEN")
)

# Apply DJL-specific configurations
default_djl_configurations, _default_max_new_tokens = _get_default_djl_configurations(
self.model, self.hf_model_config, self.schema_builder
)
self.env_vars.update(default_djl_configurations)

# Configure schema builder for text generation
if "parameters" not in self.schema_builder.sample_input:
self.schema_builder.sample_input["parameters"] = {}
self.schema_builder.sample_input["parameters"]["max_new_tokens"] = _default_max_new_tokens
# Set DJL serving defaults

# Set DJL serving defaults (only if not already set by user)
djl_env_vars = {
"OPTION_ENGINE": "Python",
"SERVING_MIN_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"OPTION_MODEL_LOADING_TIMEOUT": "240",
"OPTION_PREDICT_TIMEOUT": "60",
"TENSOR_PARALLEL_DEGREE": "1" # Default, will be overridden below
"TENSOR_PARALLEL_DEGREE": "1",
}

# Add HuggingFace authentication
if self.env_vars.get("HUGGING_FACE_HUB_TOKEN"):
djl_env_vars["HF_TOKEN"] = self.env_vars.get("HUGGING_FACE_HUB_TOKEN")

# Update with defaults only if not already set
for key, value in djl_env_vars.items():
self.env_vars.setdefault(key, value)

# DJL downloads models directly from HuggingFace Hub
self.s3_upload_path = None

Expand All@@ -367,6 +367,12 @@ def _build_for_djl(self) -> Model:
else:
self.s3_model_data_url, _ = self._prepare_for_mode()

# Set HF cache env vars to writable location (unconditionally, using setdefault
# to preserve user-provided values). This is needed because /opt/ml/model/ may be
# read-only when source_code artifacts are mounted there.
self.env_vars.setdefault("HF_HOME", "/tmp")
self.env_vars.setdefault("HUGGINGFACE_HUB_CACHE", "/tmp")

# Cache management based on mode
if self.mode in LOCAL_MODES:
self.env_vars.update({"HF_HUB_OFFLINE": "1"})
Expand Down
Empty file.
151 changes: 151 additions & 0 deletions sagemaker-serve/tests/unit/servers/test_djl_hf_cache_env.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,151 @@
"""Tests for DJL builder HF cache environment variables and HF_MODEL_ID handling.
Comment thread
aviruthen marked this conversation as resolved.

Verifies that _build_for_djl() correctly:
- Sets HF_HOME and HUGGINGFACE_HUB_CACHE to /tmp for writable cache
- Preserves user-provided HF_MODEL_ID values (uses setdefault)
- Sets HF_MODEL_ID from model param when not provided by user
- Preserves user-provided HF_HOME and HUGGINGFACE_HUB_CACHE values
"""

import pytest
from unittest.mock import Mock, patch

from sagemaker.serve.model_builder import ModelBuilder
from sagemaker.serve.utils.types import ModelServer
from sagemaker.serve.mode.function_pointers import Mode
from sagemaker.core.resources import Model


MOCK_ROLE_ARN = "arn:aws:iam::000000000000:role/SageMakerRole"
MOCK_IMAGE_URI = "000000000000.dkr.ecr.us-east-1.amazonaws.com/djl-inference:latest"
MOCK_HF_MODEL_CONFIG = {"model_type": "gpt2", "architectures": ["GPT2LMHeadModel"]}


# Common patches needed for _build_for_djl
_DJL_PATCHES = [
"sagemaker.serve.model_builder_servers._get_nb_instance",
"sagemaker.serve.model_builder_servers._get_default_djl_configurations",
"sagemaker.serve.model_builder_servers._get_model_config_properties_from_hf",
"sagemaker.serve.model_builder.ModelBuilder._is_jumpstart_model_id",
"sagemaker.serve.model_builder.ModelBuilder._validate_djl_serving_sample_data",
"sagemaker.serve.model_builder.ModelBuilder._auto_detect_image_uri",
"sagemaker.serve.model_builder.ModelBuilder._prepare_for_mode",
"sagemaker.serve.model_builder.ModelBuilder._create_model",
"sagemaker.serve.model_builder_servers._get_default_tensor_parallel_degree",
"sagemaker.serve.model_builder_servers._get_gpu_info",
]


def _mock_sagemaker_session():
"""Create a mock SageMaker session."""
session = Mock()
session.boto_region_name = "us-east-1"
session.sagemaker_config = {}
session.default_bucket.return_value = "mock-bucket"
session.upload_data.return_value = "s3://mock-bucket/model.tar.gz"
return session


def _create_djl_builder(tmp_path, env_vars=None, mode=Mode.SAGEMAKER_ENDPOINT):
"""Create a ModelBuilder configured for DJL serving tests."""
builder = ModelBuilder(
model="test-org/test-model",
role_arn=MOCK_ROLE_ARN,
sagemaker_session=_mock_sagemaker_session(),
model_path=str(tmp_path),
mode=mode,
image_uri=MOCK_IMAGE_URI,
model_server=ModelServer.DJL_SERVING,
instance_type="ml.g6e.12xlarge",
env_vars=env_vars or {},
)
builder.schema_builder = Mock()
builder.schema_builder.sample_input = {"inputs": "Hello"}
builder._optimizing = False
builder.hf_model_config = MOCK_HF_MODEL_CONFIG
return builder


def _setup_mocks(mocks):
"""Configure common mock return values for DJL build."""
# mocks are in reverse order of _DJL_PATCHES
mock_gpu_info = mocks[-1]
mock_tp_degree = mocks[-2]
mock_create = mocks[-3]
mock_prepare = mocks[-4]
# mock_auto_detect = mocks[-5] # no setup needed
# mock_validate = mocks[-6] # no setup needed
mock_is_js = mocks[-7]
mock_hf_config = mocks[-8]
mock_djl_config = mocks[-9]
mock_nb = mocks[-10]

mock_nb.return_value = None
mock_djl_config.return_value = ({}, 256)
mock_hf_config.return_value = MOCK_HF_MODEL_CONFIG
mock_is_js.return_value = False
mock_prepare.return_value = ("s3://bucket/model", None)
mock_create.return_value = Mock(spec=Model)
mock_tp_degree.return_value = 4
mock_gpu_info.return_value = 4


class TestDjlHfCacheAndModelId:
"""Tests for DJL builder HF cache env vars and HF_MODEL_ID handling."""

@pytest.fixture(autouse=True)
def _patch_djl(self):
"""Apply all DJL-related patches for each test."""
patchers = [patch(p) for p in _DJL_PATCHES]
self._mocks = [p.start() for p in patchers]
_setup_mocks(self._mocks)
yield
for p in patchers:
p.stop()

def test_sets_hf_cache_env_vars_to_tmp(self, tmp_path):
"""HF_HOME and HUGGINGFACE_HUB_CACHE should be /tmp in endpoint mode."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/tmp"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/tmp"

def test_preserves_user_provided_hf_model_id(self, tmp_path):
"""User-provided HF_MODEL_ID must NOT be overridden by model param."""
builder = _create_djl_builder(
tmp_path, env_vars={"HF_MODEL_ID": "/opt/ml/model"}
)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "/opt/ml/model"

def test_sets_hf_model_id_from_model_param_when_not_provided(self, tmp_path):
"""When no user-provided HF_MODEL_ID, it should come from model param."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "test-org/test-model"

def test_preserves_user_provided_hf_cache_dirs(self, tmp_path):
"""User-provided HF_HOME and HUGGINGFACE_HUB_CACHE should be preserved."""
builder = _create_djl_builder(
tmp_path,
env_vars={
"HF_HOME": "/my/custom/cache",
"HUGGINGFACE_HUB_CACHE": "/my/custom/hub",
},
)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/my/custom/cache"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/my/custom/hub"

def test_local_mode_sets_hf_hub_offline(self, tmp_path):
"""HF_HUB_OFFLINE=1 should be set in LOCAL_CONTAINER mode."""
builder = _create_djl_builder(tmp_path, mode=Mode.LOCAL_CONTAINER)
# Local mode doesn't need GPU info mocks for instance_type validation
builder.instance_type = None
builder._build_for_djl()

assert builder.env_vars["HF_HUB_OFFLINE"] == "1"
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
Skip to content
Merged
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
30 changes: 18 additions & 12 deletions sagemaker-serve/src/sagemaker/serve/model_builder_servers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -319,43 +319,43 @@ def _build_for_djl(self) -> Model:
logger.debug(f"Using detected notebook instance type: {nb_instance}")

if isinstance(self.model, str) and not self._is_jumpstart_model_id():
# Configure HuggingFace model for DJL
self.env_vars.update({"HF_MODEL_ID": self.model})
# Configure HuggingFace model for DJL (preserve user-provided HF_MODEL_ID)
self.env_vars.setdefault("HF_MODEL_ID", self.model)
Comment thread
aviruthen marked this conversation as resolved.

# Get model configuration for DJL optimization
self.hf_model_config = _get_model_config_properties_from_hf(
self.model, self.env_vars.get("HUGGING_FACE_HUB_TOKEN")
)

# Apply DJL-specific configurations
default_djl_configurations, _default_max_new_tokens = _get_default_djl_configurations(
self.model, self.hf_model_config, self.schema_builder
)
self.env_vars.update(default_djl_configurations)

# Configure schema builder for text generation
if "parameters" not in self.schema_builder.sample_input:
self.schema_builder.sample_input["parameters"] = {}
self.schema_builder.sample_input["parameters"]["max_new_tokens"] = _default_max_new_tokens
# Set DJL serving defaults

# Set DJL serving defaults (only if not already set by user)
djl_env_vars = {
"OPTION_ENGINE": "Python",
"SERVING_MIN_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"SERVING_MAX_WORKERS": "1",
"OPTION_MODEL_LOADING_TIMEOUT": "240",
"OPTION_PREDICT_TIMEOUT": "60",
"TENSOR_PARALLEL_DEGREE": "1" # Default, will be overridden below
"TENSOR_PARALLEL_DEGREE": "1",
}

# Add HuggingFace authentication
if self.env_vars.get("HUGGING_FACE_HUB_TOKEN"):
djl_env_vars["HF_TOKEN"] = self.env_vars.get("HUGGING_FACE_HUB_TOKEN")

# Update with defaults only if not already set
for key, value in djl_env_vars.items():
self.env_vars.setdefault(key, value)

# DJL downloads models directly from HuggingFace Hub
self.s3_upload_path = None

Expand All@@ -367,6 +367,12 @@ def _build_for_djl(self) -> Model:
else:
self.s3_model_data_url, _ = self._prepare_for_mode()

# Set HF cache env vars to writable location (unconditionally, using setdefault
# to preserve user-provided values). This is needed because /opt/ml/model/ may be
# read-only when source_code artifacts are mounted there.
self.env_vars.setdefault("HF_HOME", "/tmp")
self.env_vars.setdefault("HUGGINGFACE_HUB_CACHE", "/tmp")

# Cache management based on mode
if self.mode in LOCAL_MODES:
self.env_vars.update({"HF_HUB_OFFLINE": "1"})
Expand Down
Empty file.
151 changes: 151 additions & 0 deletions sagemaker-serve/tests/unit/servers/test_djl_hf_cache_env.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,151 @@
"""Tests for DJL builder HF cache environment variables and HF_MODEL_ID handling.
Comment thread
aviruthen marked this conversation as resolved.

Verifies that _build_for_djl() correctly:
- Sets HF_HOME and HUGGINGFACE_HUB_CACHE to /tmp for writable cache
- Preserves user-provided HF_MODEL_ID values (uses setdefault)
- Sets HF_MODEL_ID from model param when not provided by user
- Preserves user-provided HF_HOME and HUGGINGFACE_HUB_CACHE values
"""

import pytest
from unittest.mock import Mock, patch

from sagemaker.serve.model_builder import ModelBuilder
from sagemaker.serve.utils.types import ModelServer
from sagemaker.serve.mode.function_pointers import Mode
from sagemaker.core.resources import Model


MOCK_ROLE_ARN = "arn:aws:iam::000000000000:role/SageMakerRole"
MOCK_IMAGE_URI = "000000000000.dkr.ecr.us-east-1.amazonaws.com/djl-inference:latest"
MOCK_HF_MODEL_CONFIG = {"model_type": "gpt2", "architectures": ["GPT2LMHeadModel"]}


# Common patches needed for _build_for_djl
_DJL_PATCHES = [
"sagemaker.serve.model_builder_servers._get_nb_instance",
"sagemaker.serve.model_builder_servers._get_default_djl_configurations",
"sagemaker.serve.model_builder_servers._get_model_config_properties_from_hf",
"sagemaker.serve.model_builder.ModelBuilder._is_jumpstart_model_id",
"sagemaker.serve.model_builder.ModelBuilder._validate_djl_serving_sample_data",
"sagemaker.serve.model_builder.ModelBuilder._auto_detect_image_uri",
"sagemaker.serve.model_builder.ModelBuilder._prepare_for_mode",
"sagemaker.serve.model_builder.ModelBuilder._create_model",
"sagemaker.serve.model_builder_servers._get_default_tensor_parallel_degree",
"sagemaker.serve.model_builder_servers._get_gpu_info",
]


def _mock_sagemaker_session():
"""Create a mock SageMaker session."""
session = Mock()
session.boto_region_name = "us-east-1"
session.sagemaker_config = {}
session.default_bucket.return_value = "mock-bucket"
session.upload_data.return_value = "s3://mock-bucket/model.tar.gz"
return session


def _create_djl_builder(tmp_path, env_vars=None, mode=Mode.SAGEMAKER_ENDPOINT):
"""Create a ModelBuilder configured for DJL serving tests."""
builder = ModelBuilder(
model="test-org/test-model",
role_arn=MOCK_ROLE_ARN,
sagemaker_session=_mock_sagemaker_session(),
model_path=str(tmp_path),
mode=mode,
image_uri=MOCK_IMAGE_URI,
model_server=ModelServer.DJL_SERVING,
instance_type="ml.g6e.12xlarge",
env_vars=env_vars or {},
)
builder.schema_builder = Mock()
builder.schema_builder.sample_input = {"inputs": "Hello"}
builder._optimizing = False
builder.hf_model_config = MOCK_HF_MODEL_CONFIG
return builder


def _setup_mocks(mocks):
"""Configure common mock return values for DJL build."""
# mocks are in reverse order of _DJL_PATCHES
mock_gpu_info = mocks[-1]
mock_tp_degree = mocks[-2]
mock_create = mocks[-3]
mock_prepare = mocks[-4]
# mock_auto_detect = mocks[-5] # no setup needed
# mock_validate = mocks[-6] # no setup needed
mock_is_js = mocks[-7]
mock_hf_config = mocks[-8]
mock_djl_config = mocks[-9]
mock_nb = mocks[-10]

mock_nb.return_value = None
mock_djl_config.return_value = ({}, 256)
mock_hf_config.return_value = MOCK_HF_MODEL_CONFIG
mock_is_js.return_value = False
mock_prepare.return_value = ("s3://bucket/model", None)
mock_create.return_value = Mock(spec=Model)
mock_tp_degree.return_value = 4
mock_gpu_info.return_value = 4


class TestDjlHfCacheAndModelId:
"""Tests for DJL builder HF cache env vars and HF_MODEL_ID handling."""

@pytest.fixture(autouse=True)
def _patch_djl(self):
"""Apply all DJL-related patches for each test."""
patchers = [patch(p) for p in _DJL_PATCHES]
self._mocks = [p.start() for p in patchers]
_setup_mocks(self._mocks)
yield
for p in patchers:
p.stop()

def test_sets_hf_cache_env_vars_to_tmp(self, tmp_path):
"""HF_HOME and HUGGINGFACE_HUB_CACHE should be /tmp in endpoint mode."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/tmp"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/tmp"

def test_preserves_user_provided_hf_model_id(self, tmp_path):
"""User-provided HF_MODEL_ID must NOT be overridden by model param."""
builder = _create_djl_builder(
tmp_path, env_vars={"HF_MODEL_ID": "/opt/ml/model"}
)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "/opt/ml/model"

def test_sets_hf_model_id_from_model_param_when_not_provided(self, tmp_path):
"""When no user-provided HF_MODEL_ID, it should come from model param."""
builder = _create_djl_builder(tmp_path)
builder._build_for_djl()

assert builder.env_vars["HF_MODEL_ID"] == "test-org/test-model"

def test_preserves_user_provided_hf_cache_dirs(self, tmp_path):
"""User-provided HF_HOME and HUGGINGFACE_HUB_CACHE should be preserved."""
builder = _create_djl_builder(
tmp_path,
env_vars={
"HF_HOME": "/my/custom/cache",
"HUGGINGFACE_HUB_CACHE": "/my/custom/hub",
},
)
builder._build_for_djl()

assert builder.env_vars["HF_HOME"] == "/my/custom/cache"
assert builder.env_vars["HUGGINGFACE_HUB_CACHE"] == "/my/custom/hub"

def test_local_mode_sets_hf_hub_offline(self, tmp_path):
"""HF_HUB_OFFLINE=1 should be set in LOCAL_CONTAINER mode."""
builder = _create_djl_builder(tmp_path, mode=Mode.LOCAL_CONTAINER)
# Local mode doesn't need GPU info mocks for instance_type validation
builder.instance_type = None
builder._build_for_djl()

assert builder.env_vars["HF_HUB_OFFLINE"] == "1"
Loading