API Reference

ProcessingStage

View as Markdown

The ProcessingStage class is the base class for all data processing stages in NeMo Curator. Each stage defines a single step in a data curation pipeline.

Import

from nemo_curator.stages.base import ProcessingStage

Class Definition

from dataclasses import dataclass
from typing import Generic, TypeVar
InputT = TypeVar("InputT", bound=Task)
OutputT = TypeVar("OutputT", bound=Task)
@dataclass
class ProcessingStage(Generic[InputT, OutputT]):
"""Base class for all processing stages.
Type Parameters:
InputT: The input task type this stage accepts.
OutputT: The output task type this stage produces.
Class Attributes:
name: String identifier for the stage.
resources: Resources configuration (CPUs, GPUs).
batch_size: Number of tasks to process per batch.
"""
name: str = "ProcessingStage"
resources: Resources = field(default_factory=lambda: Resources(cpus=1.0))
batch_size: int = 1

Abstract Methods

inputs()

Define stage input requirements.

def inputs(self) -> tuple[list[str], list[str]]:
"""Define required task and data attributes.
Returns:
Tuple of (required_task_attributes, required_data_attributes).
"""

outputs()

Define stage output requirements.

def outputs(self) -> tuple[list[str], list[str]]:
"""Define output task and data attributes.
Returns:
Tuple of (output_task_attributes, output_data_attributes).
"""

process()

Process a single task.

def process(self, task: InputT) -> OutputT | list[OutputT] | None:
"""Process a single task.
Args:
task: The input task to process.
Returns:
- Single task: For 1-to-1 transformations
- List of tasks: For splitting/reading operations
- None: To filter out the task
"""

Optional Lifecycle Methods

setup_on_node()

Node-level initialization (e.g., download models).

def setup_on_node(
self,
node_info: NodeInfo,
worker_metadata: dict[str, Any],
) -> None:
"""Initialize resources on a compute node.
Called once per node before any workers start.
"""

setup()

Worker-level initialization (e.g., load models).

def setup(self, worker_metadata: dict[str, Any]) -> None:
"""Initialize resources for a worker.
Called once per worker before processing begins.
"""

teardown()

Cleanup after processing.

def teardown(self) -> None:
"""Clean up resources after processing completes."""

process_batch()

Vectorized batch processing for better performance.

def process_batch(self, tasks: list[InputT]) -> list[OutputT | None]:
"""Process a batch of tasks.
Override for vectorized operations.
Args:
tasks: List of input tasks.
Returns:
List of output tasks (None entries are filtered out).
"""

Backend Configuration Hooks

num_workers()

Return a backend-neutral worker count. None delegates worker sizing to the executor.

def num_workers(self) -> int | None:
return None

num_workers is reserved as a method. A subclass that defines it as a class attribute or dataclass field raises TypeError. Override the method for a class-level default, or use stage.with_(num_workers=...) for one pipeline instance.

The exact meaning depends on the executor: Ray Data creates a fixed actor or task pool, Xenna treats it as a cluster-wide count, and Ray Actor Pool caps it to available resource capacity when necessary. See Stage Worker Sizing for the complete backend matrix.

ray_stage_spec()

Return Ray-specific stage options. Ray Data consumes the worker-pool keys below, while selected flags are also used by Ray Actor Pool. Use the RayStageSpecKeys enum rather than spelling keys manually:

from nemo_curator.backends.utils import RayStageSpecKeys
def ray_stage_spec(self) -> dict:
return {
RayStageSpecKeys.IS_ACTOR_STAGE: True,
RayStageSpecKeys.MIN_WORKERS: 2,
RayStageSpecKeys.MAX_WORKERS: 8,
RayStageSpecKeys.INITIAL_WORKERS: 4,
}

xenna_stage_spec()

Return Xenna-specific stage options:

def xenna_stage_spec(self) -> dict:
return {
"num_workers_per_node": 2,
"worker_max_lifetime_m": 45,
}

Do not put num_workers in this dictionary. Use the common num_workers() hook for a cluster-wide Xenna count. num_workers() and num_workers_per_node cannot be set together.

task_id is framework-owned. The executor adapter assigns it after either process() or process_batch() returns, so custom stages must not set or derive IDs. For deterministic lineage, preserve positional correspondence in batched code: return one task or None for every input. A batch that maps multiple inputs to a different number of outputs receives random r-prefixed IDs because parentage is ambiguous.

For source-level checkpointing and the complete mapping rules, refer to Resumable Processing.

Creating Custom Stages

from dataclasses import dataclass
from nemo_curator.stages.base import ProcessingStage
from nemo_curator.stages.resources import Resources
from nemo_curator.tasks import DocumentBatch
@dataclass
class MyCustomStage(ProcessingStage[DocumentBatch, DocumentBatch]):
"""Custom stage that processes documents."""
name: str = "MyCustomStage"
resources: Resources = field(default_factory=lambda: Resources(cpus=2.0))
# Custom parameters
threshold: float = 0.5
def inputs(self) -> tuple[list[str], list[str]]:
return ["data"], ["text"]
def outputs(self) -> tuple[list[str], list[str]]:
return ["data"], ["text", "score"]
def process(self, task: DocumentBatch) -> DocumentBatch | None:
# Process the task
df = task.data
df["score"] = df["text"].apply(self._compute_score)
# Filter based on threshold
if df["score"].mean() < self.threshold:
return None
return DocumentBatch(
dataset_name=task.dataset_name,
data=df,
_metadata=task._metadata,
_stage_perf=task._stage_perf,
)
def _compute_score(self, text: str) -> float:
# Custom scoring logic
return len(text) / 1000.0

Per-Stage Runtime Environments

Stages can declare isolated Python dependencies using Ray’s native runtime_env. Set runtime_env as a class variable to specify packages that should be installed in an isolated virtualenv for that stage’s workers:

from typing import Any, ClassVar
class IsolatedStage(ProcessingStage[DocumentBatch, DocumentBatch]):
name = "isolated_stage"
runtime_env: ClassVar[dict[str, Any] | None] = {"pip": ["transformers==4.40.0"]}
def inputs(self):
return ["data"], []
def outputs(self):
return ["data"], []
def process(self, task):
import transformers # sees 4.40.0
...

You can also override runtime_env at instantiation time using with_():

stage = IsolatedStage().with_(runtime_env={"pip": ["transformers==4.45.0"]})

All three execution backends (XennaExecutor, RayDataExecutor, RayActorPoolExecutor) support per-stage runtime environments. See the Per-Stage Runtime Environments reference for details.

Configuration with with_()

with_() deep-copies a stage and configures the copy without mutating the original. It supports portable properties and backend-specific overrides:

from nemo_curator.backends.utils import RayStageSpecKeys
from nemo_curator.stages.resources import Resources
stage = MyCustomStage(threshold=0.7)
configured_stage = stage.with_(
name="configured_stage",
resources=Resources(cpus=4.0, gpus=1.0),
batch_size=8,
runtime_env={"pip": ["transformers==4.45.0"]},
num_workers=4,
# Set the stage spec for the executor you actually run with:
ray_stage_spec={
RayStageSpecKeys.RAY_NUM_CPUS: 1.0,
},
)

Both specs are shown together here only to document the arguments. In practice you set the one matching the executor you run the pipeline with — ray_stage_spec for Ray Data, xenna_stage_spec for Xenna:

# Xenna equivalent
configured_stage = stage.with_(
num_workers=4,
xenna_stage_spec={
"worker_max_lifetime_m": 45,
},
)

Setting both is harmless — each executor reads only its own spec — but it is only useful if the same stage runs under both backends.

ray_stage_spec and xenna_stage_spec are shallow-merged: user-provided top-level keys win, while nested dictionaries are replaced rather than recursively merged. An explicit None inside a stage-spec dictionary is retained. In contrast, passing the whole ray_stage_spec=None or xenna_stage_spec=None means no override.

num_workers uses an unset sentinel internally, so omission and explicit None differ:

  • Omitting num_workers preserves the stage’s current method result.
  • with_(num_workers=None) resets an inherited fixed count to executor-controlled behavior.

See Stage Worker Sizing for merge examples, precedence rules, and invalid combinations. See Per-Stage Runtime Environments for dependency isolation.

Source Code

View source on GitHub