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
Original file line number Diff line number Diff line change
Expand Up @@ -5,14 +5,14 @@

import logging
from functools import cached_property
from typing import TYPE_CHECKING

import numpy as np
import numpy.typing as npt
import pandas as pd
from numpy.linalg import norm
from pydantic import BaseModel, ConfigDict, Field
from scipy.stats import ks_2samp
from sentence_transformers import SentenceTransformer
from tenacity import (
RetryError,
Retrying,
Expand All @@ -35,6 +35,9 @@
from ...observability import get_logger
from . import multi_modal_figures as figures

if TYPE_CHECKING:
from sentence_transformers import SentenceTransformer

logger = get_logger(__name__)


Expand Down Expand Up @@ -237,6 +240,11 @@ def _preprocess_text_data(text_data: pd.Series, nrows: int) -> pd.Series:
@staticmethod
def _init_sentence_transformer_model() -> SentenceTransformer | None:
"""Load the sentence transformer model with exponential-backoff retries."""
try:
from sentence_transformers import SentenceTransformer
except ImportError:
return None

try:
for attempt in Retrying(
# TODO(PLAT-2537): Temporarily increase retries, but we will bundle this model next, so download is not required.
Expand Down
8 changes: 8 additions & 0 deletions tests/e2e/test_safe_synthesizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,12 @@
logger = get_logger(__name__)


def _assert_evaluation_report_rendered(report_html: str | None) -> None:
assert report_html is not None
assert "Synthetic Quality Score" in report_html
assert "Text Semantic Similarity" in report_html


@pytest.mark.e2e
@pytest.mark.requires_gpu
@pytest.mark.timeout(1000)
Expand All @@ -66,6 +72,7 @@ def test_train_and_generate_dp(fixture_financial_transactions_dataset, fixture_s
assert result.summary.timing.training_time_sec is not None and result.summary.timing.training_time_sec > 0
assert result.summary.timing.generation_time_sec is not None and result.summary.timing.generation_time_sec > 0
assert result.summary.timing.evaluation_time_sec is not None and result.summary.timing.evaluation_time_sec > 0
_assert_evaluation_report_rendered(result.evaluation_report_html)


@pytest.mark.e2e
Expand All @@ -89,3 +96,4 @@ def test_train_and_generate_defaults(fixture_financial_transactions_dataset, fix
assert result.summary.timing.training_time_sec is not None and result.summary.timing.training_time_sec > 0
assert result.summary.timing.generation_time_sec is not None and result.summary.timing.generation_time_sec > 0
assert result.summary.timing.evaluation_time_sec is not None and result.summary.timing.evaluation_time_sec > 0
_assert_evaluation_report_rendered(result.evaluation_report_html)
52 changes: 32 additions & 20 deletions tests/evaluation/reports/test_multimodal_report.py
Original file line number Diff line number Diff line change
@@ -1,22 +1,36 @@
# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0

import json

import pandas as pd
import pytest

# Skip all tests in this module if sentence_transformers is not available
pytest.importorskip(
"sentence_transformers",
reason="sentence_transformers is required for these tests (install with: uv sync --extra cpu)",
)

from nemo_safe_synthesizer.config.evaluate import EvaluationParameters
from nemo_safe_synthesizer.config.parameters import SafeSynthesizerParameters
from nemo_safe_synthesizer.evaluation.components.text_semantic_similarity import TextSemanticSimilarity
from nemo_safe_synthesizer.evaluation.data_model.evaluation_datasets import EvaluationDatasets
from nemo_safe_synthesizer.evaluation.data_model.evaluation_score import Grade
from nemo_safe_synthesizer.evaluation.data_model.evaluation_score import EvaluationScore, Grade
from nemo_safe_synthesizer.evaluation.reports.multimodal import multimodal_report as multimodal_report_module
from nemo_safe_synthesizer.evaluation.reports.multimodal.multimodal_report import MultimodalReport


@pytest.fixture(autouse=True)
def stub_text_semantic_similarity(monkeypatch: pytest.MonkeyPatch) -> None:
"""Keep report assembly tests independent of sentence-transformer inference."""

class StubTextSemanticSimilarity:
Comment thread
mckornfield marked this conversation as resolved.
@staticmethod
def from_evaluation_datasets(*_args, **_kwargs) -> TextSemanticSimilarity:
return TextSemanticSimilarity(score=EvaluationScore.finalize_grade(9.0, 9.0))

monkeypatch.setattr(
multimodal_report_module,
"TextSemanticSimilarity",
StubTextSemanticSimilarity,
)
Comment thread
mckornfield marked this conversation as resolved.


def _minimal_multimodal_report() -> MultimodalReport:
training_df = pd.DataFrame({"x": [1, 2], "y": [3, 4]})
synthetic_df = pd.DataFrame({"x": [1, 2], "y": [3, 4]})
Expand All @@ -38,7 +52,6 @@ def test_jinja_context_job_id_set_when_nemo_job_id_present(monkeypatch: pytest.M
assert ctx["job_id"] == "cluster-job-abc123"


@pytest.mark.slow
def test_from_dataframes_applies_sqs_report_config(fixture_training_df, fixture_synthetic_df, fixture_test_df) -> None:
"""``sqs_report_rows`` / ``sqs_report_columns`` from config drive the actual subsampling.

Expand Down Expand Up @@ -72,13 +85,12 @@ def test_from_dataframes_applies_sqs_report_config(fixture_training_df, fixture_
assert report.evaluation_datasets.synthetic_cols == target_cols


@pytest.mark.slow
def test_multimodal_report(
fixture_training_df_5k, fixture_synthetic_df_5k, fixture_test_df, fixture_skip_privacy_metrics_config
fixture_training_df, fixture_synthetic_df, fixture_test_df, fixture_skip_privacy_metrics_config
):
report = MultimodalReport.from_dataframes(
training=fixture_training_df_5k,
synthetic=fixture_synthetic_df_5k,
training=fixture_training_df,
synthetic=fixture_synthetic_df,
test=fixture_test_df,
config=fixture_skip_privacy_metrics_config,
)
Expand All @@ -89,15 +101,15 @@ def test_multimodal_report(

report_dict = report.get_dict()
assert len(report_dict) == 6
assert report_dict["Synthetic Quality Score"] == {
"raw_score": 9.7935,
assert report_dict["Text Semantic Similarity"] == {
"raw_score": 9.0,
"grade": "Excellent",
"score": 9.8,
"score": 9.0,
"notes": None,
}
assert report_dict["Synthetic Quality Score"]["grade"] == "Excellent"
assert report_dict["Synthetic Quality Score"]["score"] > 0

report_json = report.get_json()
assert (
'"Synthetic Quality Score": {"raw_score": 9.7935, "grade": "Excellent", "score": 9.8, "notes": null}}'
in report_json
)
report_json = json.loads(report.get_json())
assert report_json["Text Semantic Similarity"] == report_dict["Text Semantic Similarity"]
assert report_json["Synthetic Quality Score"] == report_dict["Synthetic Quality Score"]
Loading