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
63 changes: 49 additions & 14 deletions nemo_automodel/_transformers/sentence_transformer_export.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,9 +64,15 @@
"normalize": {
"sentence_transformers.models.Normalize",
"sentence_transformers.sentence_transformer.modules.normalize.Normalize",
"sentence_transformers.base.modules.normalize.Normalize",
},
}
_SUPPORTED_SENTENCE_TRANSFORMER_MODULE_STACKS = {
("transformer", "pooling"),
("transformer", "pooling", "normalize"),
}
_SENTENCE_TRANSFORMER_EXPORT_MODULE_TYPES = {
# v6 remaps these v5.4-era paths, keeping exported checkpoints loadable across v5.4+.
"transformer": "sentence_transformers.base.modules.transformer.Transformer",
"pooling": "sentence_transformers.sentence_transformer.modules.pooling.Pooling",
"normalize": "sentence_transformers.sentence_transformer.modules.normalize.Normalize",
Expand Down Expand Up @@ -154,24 +160,53 @@ def _load_sentence_transformer_wrapper_options(
if not isinstance(modules, list):
raise ValueError("Sentence Transformers modules.json must contain a list of modules.")

expected_module_types = [
_SENTENCE_TRANSFORMER_MODULE_TYPES["transformer"],
_SENTENCE_TRANSFORMER_MODULE_TYPES["pooling"],
]
module_types = [module.get("type") if isinstance(module, dict) else None for module in modules]
if len(module_types) == 3:
expected_module_types.append(_SENTENCE_TRANSFORMER_MODULE_TYPES["normalize"])
if len(module_types) not in (2, 3) or any(
module_type not in allowed_types
for module_type, allowed_types in zip(module_types, expected_module_types, strict=True)
):
module_roles = []
for module in modules:
module_type = module.get("type") if isinstance(module, dict) else None
module_role = next(
(
role
for role, allowed_types in _SENTENCE_TRANSFORMER_MODULE_TYPES.items()
if module_type in allowed_types
),
None,
)
module_roles.append(module_role)
module_roles = tuple(module_roles)
if module_roles not in _SUPPORTED_SENTENCE_TRANSFORMER_MODULE_STACKS:
raise ValueError(
"Sentence Transformers metadata must use the exact supported module stack: "
"Transformer, Pooling, and optional Normalize."
)
if modules[0].get("path") != "":

modules_by_role = dict(zip(module_roles, modules, strict=True))
transformer_module = modules_by_role["transformer"]
pooling_module = modules_by_role["pooling"]
normalize_module = modules_by_role.get("normalize")
if transformer_module.get("path") != "":
raise ValueError("Sentence Transformers Transformer metadata must reference the checkpoint root.")

if normalize_module is not None:
normalize_path = normalize_module.get("path")
if not isinstance(normalize_path, str) or not normalize_path:
raise ValueError("Sentence Transformers Normalize metadata must reference a module path.")
normalize_config = _load_sentence_transformer_json(
model_name_or_path,
os.path.join(normalize_path, "config.json"),
hf_kwargs,
)
if normalize_config is not None:
if not isinstance(normalize_config, dict):
raise ValueError("Sentence Transformers Normalize config.json is invalid.")
normalize_input_name = normalize_config.get("module_input_name", "sentence_embedding")
normalize_output_name = normalize_config.get("module_output_name")
if normalize_output_name is None:
normalize_output_name = normalize_input_name
if normalize_input_name != "sentence_embedding" or normalize_output_name != "sentence_embedding":
raise ValueError(
"Sentence Transformers Normalize metadata must normalize the final sentence embedding in place."
)

sentence_bert_config = _load_sentence_transformer_json(
model_name_or_path,
"sentence_bert_config.json",
Expand All @@ -183,7 +218,7 @@ def _load_sentence_transformer_wrapper_options(
"the NeMo inference path does not lowercase text."
)

pooling_path = modules[1].get("path")
pooling_path = pooling_module.get("path")
if not isinstance(pooling_path, str) or not pooling_path:
raise ValueError("Sentence Transformers Pooling metadata must reference a module path.")
pooling_config = _load_sentence_transformer_json(
Expand Down Expand Up @@ -246,7 +281,7 @@ def _load_sentence_transformer_wrapper_options(

return SentenceTransformerWrapperOptions(
pooling=matching_pooling[0],
l2_normalize=len(modules) == 3,
l2_normalize=normalize_module is not None,
query_prompt=query_prompt,
document_prompt=document_prompt,
)
Expand Down
113 changes: 107 additions & 6 deletions tests/unit_tests/_transformers/test_sentence_transformer_export.py
Original file line number Diff line number Diff line change
Expand Up @@ -371,18 +371,34 @@ def test_cache_hub_source_legal_assets_uses_exact_loaded_revision(monkeypatch):
)


def test_sentence_transformer_source_with_unsupported_module_is_rejected(tmp_path):
(tmp_path / "1_Pooling").mkdir()
@pytest.mark.parametrize(
"module_types",
[
[
"sentence_transformers.models.Transformer",
"sentence_transformers.models.Pooling",
"sentence_transformers.models.Dense",
],
[
"sentence_transformers.models.Transformer",
"sentence_transformers.models.Normalize",
"sentence_transformers.models.Pooling",
],
],
)
def test_sentence_transformer_source_with_unsupported_module_stack_is_rejected(tmp_path, module_types):
(tmp_path / "modules.json").write_text(
json.dumps(
[
{"idx": 0, "path": "", "type": "sentence_transformers.models.Transformer"},
{"idx": 1, "path": "1_Pooling", "type": "sentence_transformers.models.Pooling"},
{"idx": 2, "path": "2_Dense", "type": "sentence_transformers.models.Dense"},
{
"idx": index,
"path": "" if index == 0 else f"{index}_Module",
"type": module_type,
}
for index, module_type in enumerate(module_types)
]
)
)
(tmp_path / "1_Pooling" / "config.json").write_text(json.dumps({"pooling_mode_mean_tokens": True}))

with pytest.raises(ValueError, match="exact supported module stack"):
sentence_transformer_export._load_sentence_transformer_wrapper_options(str(tmp_path), {})
Expand Down Expand Up @@ -452,6 +468,91 @@ def test_sentence_transformer_source_accepts_canonical_module_paths(tmp_path, mo
assert options.l2_normalize is expected_normalize


def test_sentence_transformer_source_accepts_v6_normalize_module(tmp_path):
(tmp_path / "1_Pooling").mkdir()
(tmp_path / "2_Normalize").mkdir()
(tmp_path / "modules.json").write_text(
json.dumps(
[
{
"idx": 0,
"path": "",
"type": "sentence_transformers.base.modules.transformer.Transformer",
},
{
"idx": 1,
"path": "1_Pooling",
"type": "sentence_transformers.sentence_transformer.modules.pooling.Pooling",
},
{
"idx": 2,
"path": "2_Normalize",
"type": "sentence_transformers.base.modules.normalize.Normalize",
},
]
)
)
(tmp_path / "1_Pooling" / "config.json").write_text(json.dumps({"pooling_mode": "mean"}))
(tmp_path / "2_Normalize" / "config.json").write_text(
json.dumps(
{
"module_input_name": "sentence_embedding",
"module_output_name": "sentence_embedding",
}
)
)

options = sentence_transformer_export._load_sentence_transformer_wrapper_options(str(tmp_path), {})

assert options is not None
assert options.pooling == "avg"
assert options.l2_normalize is True


@pytest.mark.parametrize(
"normalize_config",
[
{
"module_input_name": "token_embeddings",
"module_output_name": "token_embeddings",
},
{
"module_input_name": "sentence_embedding",
"module_output_name": "custom_embedding",
},
],
)
def test_sentence_transformer_source_rejects_unrepresentable_v6_normalize_config(tmp_path, normalize_config):
(tmp_path / "1_Pooling").mkdir()
(tmp_path / "2_Normalize").mkdir()
(tmp_path / "modules.json").write_text(
json.dumps(
[
{
"idx": 0,
"path": "",
"type": "sentence_transformers.base.modules.transformer.Transformer",
},
{
"idx": 1,
"path": "1_Pooling",
"type": "sentence_transformers.sentence_transformer.modules.pooling.Pooling",
},
{
"idx": 2,
"path": "2_Normalize",
"type": "sentence_transformers.base.modules.normalize.Normalize",
},
]
)
)
(tmp_path / "1_Pooling" / "config.json").write_text(json.dumps({"pooling_mode": "mean"}))
(tmp_path / "2_Normalize" / "config.json").write_text(json.dumps(normalize_config))

with pytest.raises(ValueError, match="final sentence embedding"):
sentence_transformer_export._load_sentence_transformer_wrapper_options(str(tmp_path), {})


@pytest.mark.parametrize(
"module_types",
[
Expand Down
7 changes: 4 additions & 3 deletions uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading