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 @@ -137,6 +137,8 @@ class FaithEvalFilter(ProcessingStage[DocumentBatch, DocumentBatch]):
downstream score analysis before committing to a threshold.
generation_config : GenerationConfig | None
LLM generation parameters. Defaults to ``temperature=0.0, max_tokens=256``.
prompt_path : str | None
Absolute local YAML prompt path, or ``None`` for the packaged prompt.
"""

name: str = "FaithEvalFilter"
Expand All @@ -149,6 +151,7 @@ class FaithEvalFilter(ProcessingStage[DocumentBatch, DocumentBatch]):
threshold: float = 2.5
filter_enabled: bool = True
generation_config: GenerationConfig | None = None
prompt_path: str | None = None
max_concurrent_requests: int = 64

# -- internal state (not constructor args) ---------------------------------
Expand Down Expand Up @@ -191,7 +194,8 @@ def setup(self, worker_metadata: WorkerMetadata | None = None) -> None: # noqa:
runs on the driver, while ``setup()`` runs on the worker.
"""
if not self._initialized:
self._system_prompt, self._user_template = load_prompt_template("faith_eval.yaml")
prompt_file = self.prompt_path or "faith_eval.yaml"

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

what if prompt_path=""

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Empty string will go to "faith_eval.yaml".

self._system_prompt, self._user_template = load_prompt_template(prompt_file)

if self.generation_config is None:
self.generation_config = GenerationConfig(
Expand Down
16 changes: 16 additions & 0 deletions nemo_curator/stages/text/experimental/translation/pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,11 @@ class TranslationStage(CompositeStage[DocumentBatch, DocumentBatch]):
client: AsyncLLMClient | None = None
model_name: str = ""
generation_config: GenerationConfig | None = None
translation_prompt_path: str | None = None
max_concurrent_requests: int = 64
health_check: bool = True
dry_run: bool = False
dry_run_log_count: int = 5

backend_type: str = "llm"
backend_config: dict = field(default_factory=dict)
Expand All @@ -67,6 +72,9 @@ class TranslationStage(CompositeStage[DocumentBatch, DocumentBatch]):
faith_threshold: float = 2.5
faith_model_name: str = ""
filter_enabled: bool = True
faith_generation_config: GenerationConfig | None = None
faith_prompt_path: str | None = None
faith_max_concurrent_requests: int = 64

output_mode: str = "replaced"
merge_scores: bool = False
Expand Down Expand Up @@ -178,6 +186,11 @@ def _build_stages(self) -> list[ProcessingStage]:
backend_type=self.backend_type,
backend_config=self.backend_config,
generation_config=self.generation_config,
prompt_path=self.translation_prompt_path,
max_concurrent_requests=self.max_concurrent_requests,
health_check=self.health_check,
dry_run=self.dry_run,
dry_run_log_count=self.dry_run_log_count,
)
)

Expand All @@ -193,6 +206,9 @@ def _build_stages(self) -> list[ProcessingStage]:
translated_text_field="_translated",
threshold=self.faith_threshold,
filter_enabled=False,
generation_config=self.faith_generation_config,
prompt_path=self.faith_prompt_path,
max_concurrent_requests=self.faith_max_concurrent_requests,
)
)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,17 @@
"faith_handling_of_format": "Handling_of_Format",
}

_OUTPUT_COLUMN_DTYPES = {
"translation_time": "float64",
"faith_fluency": "float64",
"faith_accuracy": "float64",
"faith_idiomaticity": "float64",
"faith_terminology": "float64",
"faith_handling_of_format": "float64",
"faith_avg": "float64",
"faith_parse_failed": "bool",
}


@dataclass
class ReassemblyStage(ProcessingStage[DocumentBatch, DocumentBatch]):
Expand Down Expand Up @@ -86,6 +97,21 @@ def process(self, batch: DocumentBatch) -> DocumentBatch:
"""Reassemble translated segments into full documents."""
df = batch.to_pandas()

if df.empty:
logger.info("ReassemblyStage: no translated segment rows to reassemble")
base_cols = [col for col in df.columns if col not in _INTERNAL_COLUMNS]
out_df = df.loc[:, base_cols].copy()
for col in self.outputs()[1]:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

is it a bit fragile that it requires outputs() returning a tuple where index [1] contains the output column names

@sarahyurick sarahyurick Jun 3, 2026 •

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Agree it looks a bit fragile but it follows the convention that should be used by text stages with DocumentBatch everywhere.

if col not in out_df.columns:
out_df[col] = pd.Series(dtype=_OUTPUT_COLUMN_DTYPES.get(col, "object"))
return DocumentBatch(
task_id=batch.task_id,
dataset_name=batch.dataset_name,
data=out_df,
_metadata=batch._metadata,
_stage_perf=batch._stage_perf,
)

result_rows: list[dict[str, Any]] = []

for _doc_id, doc_group in df.groupby("_seg_doc_id", sort=True):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -315,6 +315,20 @@ def process(self, batch: DocumentBatch) -> DocumentBatch:
df = batch.to_pandas()
field_paths = normalize_text_field(self.text_field)

if df.empty:
logger.info("SegmentationStage: no documents to segment")
out_df = df.copy()
out_df["_seg_segments"] = pd.Series(dtype="object")
out_df["_seg_metadata"] = pd.Series(dtype="object")
out_df["_seg_doc_id"] = pd.Series(dtype="int64")
return DocumentBatch(
task_id=batch.task_id,
dataset_name=batch.dataset_name,
data=out_df,
_metadata=batch._metadata,
_stage_perf=batch._stage_perf,
)

all_rows: list[dict[str, Any]] = []

total_docs = len(df)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,11 @@ def process(self, batch: DocumentBatch) -> DocumentBatch:
"translation_column": self.translation_column,
}

if len(df) > 0 and remaining_df.empty:
logger.warning(
"SkipExistingTranslationsStage: no rows require translation after skipping already-translated rows"
)

logger.info(
"SkipExistingTranslationsStage: skipping {} already-translated rows, processing {} rows",
len(skipped_rows),
Expand Down Expand Up @@ -150,7 +155,7 @@ def process(self, batch: DocumentBatch) -> DocumentBatch:
continue
skipped_df[col] = self._COLUMN_DEFAULTS.get(col, "")

merged = pd.concat([df, skipped_df], ignore_index=True)
merged = skipped_df.reset_index(drop=True) if df.empty else pd.concat([df, skipped_df], ignore_index=True)
if order_col in merged.columns:
merged = merged.sort_values(order_col).reset_index(drop=True)
merged = merged.drop(columns=[order_col])
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,8 @@ class SegmentTranslationStage(ProcessingStage[DocumentBatch, DocumentBatch]):
backend_type: str = "llm"
backend_config: dict = field(default_factory=dict)
generation_config: GenerationConfig | None = None
prompt_path: str | None = None
"""Absolute local YAML prompt path, or ``None`` for the packaged prompt."""
max_concurrent_requests: int = 64
health_check: bool = True
"""If True, verify the translation backend is reachable during ``setup()``."""
Expand Down Expand Up @@ -101,7 +103,8 @@ def outputs(self) -> tuple[list[str], list[str]]:
def setup(self, worker_metadata: WorkerMetadata | None = None) -> None: # noqa: ARG002
"""Initialize the client or backend on the worker."""
if not self._initialized:
self._system_prompt, self._user_template = load_prompt_template("translate.yaml")
prompt_file = self.prompt_path or "translate.yaml"
self._system_prompt, self._user_template = load_prompt_template(prompt_file)
Comment on lines +106 to +107

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 prompt_path = "" silently falls back to the default

self.prompt_path or "translate.yaml" treats an empty string as falsy, so a caller who accidentally passes prompt_path="" gets the packaged prompt with no warning. The type annotation (str | None) implies the only "no-op" value is None, so the guard should be explicit. The same pattern appears in FaithEvalFilter.setup() at the corresponding line.


if self.backend_type != "llm":
from nemo_curator.stages.text.experimental.translation.backends import get_backend
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,14 +23,15 @@
_PROMPT_DIR = Path(__file__).resolve().parent.parent / "prompts"


def load_prompt_template(filename: str) -> tuple[str, str]:
def load_prompt_template(filename_or_path: str | Path) -> tuple[str, str]:
"""Load a YAML prompt file and return ``(system_prompt, user_template)``.

Parameters
----------
filename : str
Name of the YAML file inside the ``prompts/`` directory
(e.g. ``"translate.yaml"`` or ``"faith_eval.yaml"``).
filename_or_path : str | Path
Name of a YAML file inside the packaged ``prompts/`` directory
(e.g. ``"translate.yaml"`` or ``"faith_eval.yaml"``), or an absolute
local path to a YAML prompt file.

Returns
-------
Expand All @@ -46,7 +47,9 @@ def load_prompt_template(filename: str) -> tuple[str, str]:
KeyError
If the top-level mapping is missing the ``system`` or ``user`` key.
"""
prompt_path = _PROMPT_DIR / filename
prompt_path = Path(filename_or_path)
if not prompt_path.is_absolute():
prompt_path = _PROMPT_DIR / prompt_path
try:
with open(prompt_path, encoding="utf-8") as fh:
data = yaml.safe_load(fh)
Expand Down
84 changes: 84 additions & 0 deletions tests/stages/text/experimental/translation/test_pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
from __future__ import annotations

import json
from typing import TYPE_CHECKING

import pandas as pd
import pytest
Expand All @@ -42,6 +43,9 @@
)
from nemo_curator.tasks import DocumentBatch

if TYPE_CHECKING:
from pathlib import Path

from .conftest import MockAsyncLLMClient


Expand Down Expand Up @@ -87,6 +91,27 @@ def test_decompose_faith_eval_fallback_to_model_name(self, mock_client: MockAsyn
faith_stage = _only_stage_of_type(stages, FaithEvalFilter)
assert faith_stage.model_name == "translate-model"

def test_decompose_faith_eval_inherits_control_knobs(self, mock_client: MockAsyncLLMClient) -> None:
"""FaithEvalFilter receives model, prompt, generation, and concurrency config."""
gen_cfg = GenerationConfig(temperature=0.0, max_tokens=128)
pipeline = TranslationStage(
source_lang="en",
target_lang="de",
client=mock_client,
model_name="translate-model",
enable_faith_eval=True,
faith_model_name="faith-model",
faith_generation_config=gen_cfg,
faith_prompt_path="/opt/prompts/custom_faith.yaml",
faith_max_concurrent_requests=3,
)
stages = pipeline.decompose()
faith_stage = _only_stage_of_type(stages, FaithEvalFilter)
assert faith_stage.model_name == "faith-model"
assert faith_stage.generation_config is gen_cfg
assert faith_stage.prompt_path == "/opt/prompts/custom_faith.yaml"
assert faith_stage.max_concurrent_requests == 3

def test_llm_backend_requires_model_name(self, mock_client: MockAsyncLLMClient) -> None:
"""LLM translation should fail fast when model_name is unset."""
with pytest.raises(ValueError, match="non-empty 'model_name'"):
Expand Down Expand Up @@ -197,6 +222,11 @@ def test_translate_stage_inherits_config(self, mock_client: MockAsyncLLMClient)
client=mock_client,
model_name="m",
generation_config=gen_cfg,
translation_prompt_path="/opt/prompts/custom_translate.yaml",
max_concurrent_requests=7,
health_check=False,
dry_run=True,
dry_run_log_count=2,
backend_type="llm",
)
tr = pipeline.decompose()[1]
Expand All @@ -205,6 +235,11 @@ def test_translate_stage_inherits_config(self, mock_client: MockAsyncLLMClient)
assert tr.target_lang == "ja"
assert tr.model_name == "m"
assert tr.generation_config is gen_cfg
assert tr.prompt_path == "/opt/prompts/custom_translate.yaml"
assert tr.max_concurrent_requests == 7
assert tr.health_check is False
assert tr.dry_run is True
assert tr.dry_run_log_count == 2

def test_reassembly_stage_inherits_config(self, mock_client: MockAsyncLLMClient) -> None:
"""ReassemblyStage receives text_field and output_field from the pipeline."""
Expand Down Expand Up @@ -260,6 +295,26 @@ def test_pipeline_inputs_outputs(self, mock_client: MockAsyncLLMClient) -> None:
class TestFaithEvalFilter:
"""Tests for FaithEvalFilter score parsing and filtering."""

def test_setup_loads_custom_prompt_path(self, mock_client: MockAsyncLLMClient, tmp_path: Path) -> None:
"""Verify setup() can load a caller-provided absolute FAITH prompt path."""
prompt_path = tmp_path / "custom_faith.yaml"
prompt_path.write_text(
"system: custom faith {source_language} {target_language}\nuser: custom {source_text} {translated_text}\n",
encoding="utf-8",
)
stage = FaithEvalFilter(
source_lang="en",
target_lang="de",
client=mock_client,
model_name="faith-model",
prompt_path=str(prompt_path),
)

stage.setup(worker_metadata=None)

assert stage._system_prompt == "custom faith {source_language} {target_language}"
assert stage._user_template == "custom {source_text} {translated_text}"

def test_extract_scores_valid_json(self) -> None:
"""Valid JSON with all 5 keys is parsed correctly."""
text = '{"Fluency": 4, "Accuracy": 5, "Idiomaticity": 3, "Terminology": 4, "Handling_of_Format": 5}'
Expand Down Expand Up @@ -775,6 +830,35 @@ def test_skip_translated_false_retranslates_all(
assert len(result_df) == 3
assert all(len(t) > 0 for t in result_df["translated_text"])

def test_skip_translated_all_rows_already_translated(self, mock_client: MockAsyncLLMClient) -> None:
"""An all-skipped batch should pass through and restore rows without missing-column errors."""
df = pd.DataFrame(
{
"id": [100, 200],
"text": ["Already translated", "Already translated too"],
"translated_text": ["Bereits uebersetzt", "Auch bereits uebersetzt"],
}
)
batch = DocumentBatch(data=df, dataset_name="resume-test", task_id="1")
pipeline = TranslationStage(
source_lang="en",
target_lang="de",
client=mock_client,
model_name="test-model",
skip_translated=True,
health_check=False,
)

result = batch
for stage in pipeline.decompose():
stage.setup()
result = stage.process(result)

result_df = result.to_pandas()
assert list(result_df["id"]) == [100, 200]
assert list(result_df["translated_text"]) == ["Bereits uebersetzt", "Auch bereits uebersetzt"]
assert "_skipped_rows_state" not in result._metadata

def test_merge_skipped_reads_batch_metadata(self) -> None:
"""Skipped-row state should travel with the batch, not the stage instance."""
df = pd.DataFrame(
Expand Down
22 changes: 22 additions & 0 deletions tests/stages/text/experimental/translation/test_prompts.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,12 @@
import sys
from pathlib import Path
from types import SimpleNamespace

import pytest

from nemo_curator.stages.text.experimental.translation.utils.prompt_loader import (
load_prompt_template,
)
from nemo_curator.stages.text.utils.text_utils import get_language_name


Expand Down Expand Up @@ -35,3 +39,21 @@ def test_get_language_name_falls_back_on_unknown_code(monkeypatch: pytest.Monkey
SimpleNamespace(to_name=lambda code: (_ for _ in ()).throw(KeyError(code))),
)
assert get_language_name("zz") == "zz"


def test_load_prompt_template_supports_absolute_path(tmp_path: Path) -> None:
prompt_path = tmp_path / "custom_translate.yaml"
prompt_path.write_text("system: custom system\nuser: custom user {src}\n", encoding="utf-8")

system_prompt, user_template = load_prompt_template(prompt_path)

assert system_prompt == "custom system"
assert user_template == "custom user {src}"


def test_load_prompt_template_rejects_missing_required_keys(tmp_path: Path) -> None:
prompt_path = tmp_path / "bad_prompt.yaml"
prompt_path.write_text("system: custom system\n", encoding="utf-8")

with pytest.raises(KeyError, match="missing required keys"):
load_prompt_template(prompt_path)
Loading
Loading