API ReferenceTasks
ImageBatch
ImageBatch is the task type for image processing in NeMo Curator.
Import
from nemo_curator.tasks import ImageBatch
Class Definition
from dataclasses import dataclassfrom nemo_curator.tasks.image import ImageObject@dataclassclass ImageBatch(Task[list[ImageObject]]):"""Task containing a batch of images.Attributes:task_id: Framework-managed deterministic lineage identifier.dataset_name: Name of the source dataset.data: List of ImageObject instances."""dataset_name: strdata: list[ImageObject]# task_id is inherited from Task with init=False.
ImageObject
Each image in the batch is represented by an ImageObject:
@dataclassclass ImageObject:"""Represents a single image with metadata.Attributes:path: Path to the image file.caption: Optional text caption for the image.metadata: Additional metadata dictionary.embeddings: Optional embedding vector."""path: strcaption: str | None = Nonemetadata: dict[str, Any] = field(default_factory=dict)embeddings: np.ndarray | None = None
Properties
num_items
Get the number of images in the batch.
@propertydef num_items(self) -> int:"""Returns the number of images in this batch."""
Creating ImageBatch
from nemo_curator.tasks import ImageBatchfrom nemo_curator.tasks.image import ImageObject# Create image objectsimages = [ImageObject(path="/data/images/image1.jpg",caption="A cat sitting on a couch",metadata={"source": "dataset_a"},),ImageObject(path="/data/images/image2.jpg",caption="A dog playing in the park",metadata={"source": "dataset_a"},),]# Create batchbatch = ImageBatch(dataset_name="image_dataset",data=images,)
Do not pass task_id to the constructor. The framework assigns it after the task crosses a stage boundary.
Usage in Stages
from dataclasses import dataclassfrom nemo_curator.stages.base import ProcessingStagefrom nemo_curator.tasks import ImageBatch@dataclassclass ImageFilterStage(ProcessingStage[ImageBatch, ImageBatch]):"""Filter images based on metadata."""name: str = "ImageFilter"min_resolution: int = 256def inputs(self) -> tuple[list[str], list[str]]:return ["data"], []def outputs(self) -> tuple[list[str], list[str]]:return ["data"], []def process(self, task: ImageBatch) -> ImageBatch | None:filtered = [img for img in task.dataif img.metadata.get("width", 0) >= self.min_resolutionand img.metadata.get("height", 0) >= self.min_resolution]if not filtered:return Nonereturn ImageBatch(dataset_name=task.dataset_name,data=filtered,_metadata=task._metadata,_stage_perf=task._stage_perf,)
Common Operations
Adding Embeddings
def process(self, task: ImageBatch) -> ImageBatch:for img in task.data:img.embeddings = self.model.encode(img.path)return ImageBatch(dataset_name=task.dataset_name,data=task.data,_metadata=task._metadata,_stage_perf=task._stage_perf,)
Filtering by Score
def process(self, task: ImageBatch) -> ImageBatch | None:filtered = [img for img in task.dataif img.metadata.get("aesthetic_score", 0) >= self.threshold]if not filtered:return Nonereturn ImageBatch(dataset_name=task.dataset_name,data=filtered,_metadata=task._metadata,_stage_perf=task._stage_perf,)