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
7 changes: 7 additions & 0 deletions examples/audio/metadata_extraction/run_metadata_extraction.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,7 +65,7 @@
_SORTFORMER_BATCH_WINDOW_MULTIPLIER = 4


def _build_arg_parser() -> argparse.ArgumentParser:

Check failure on line 68 in examples/audio/metadata_extraction/run_metadata_extraction.py

View workflow job for this annotation

GitHub Actions / ruff

Ruff (PLR0915)

examples/audio/metadata_extraction/run_metadata_extraction.py:68:5: PLR0915 Too many statements (68 > 50)

Check failure on line 68 in examples/audio/metadata_extraction/run_metadata_extraction.py

View workflow job for this annotation

GitHub Actions / ruff

Ruff (PLR0915)

examples/audio/metadata_extraction/run_metadata_extraction.py:68:5: PLR0915 Too many statements (68 > 50)
ap = argparse.ArgumentParser(description="Metadata extraction pipeline for unsegmented audio")

ap.add_argument("--data_config", type=str, default=None, help="Path to input_cfg YAML.")
Expand Down Expand Up @@ -131,6 +131,12 @@
"Use FLOAT to avoid quantization changes in diarization output."
),
)
ap.add_argument(
"--max_audio_duration_sec",
type=float,
default=12 * 60 * 60,
help="Maximum source-audio duration to process (default: 12 hours; 0 disables the limit).",
)

vad = ap.add_argument_group("VAD (Silero)")
vad.add_argument(
Expand Down Expand Up @@ -360,7 +366,7 @@
return ap


def _build_stages(args: argparse.Namespace, language_filter: list[str] | None) -> list:

Check failure on line 369 in examples/audio/metadata_extraction/run_metadata_extraction.py

View workflow job for this annotation

GitHub Actions / ruff

Ruff (PLR0915)

examples/audio/metadata_extraction/run_metadata_extraction.py:369:5: PLR0915 Too many statements (52 > 50)

Check failure on line 369 in examples/audio/metadata_extraction/run_metadata_extraction.py

View workflow job for this annotation

GitHub Actions / ruff

Ruff (PLR0912)

examples/audio/metadata_extraction/run_metadata_extraction.py:369:5: PLR0912 Too many branches (15 > 12)

Check failure on line 369 in examples/audio/metadata_extraction/run_metadata_extraction.py

View workflow job for this annotation

GitHub Actions / ruff

Ruff (C901)

examples/audio/metadata_extraction/run_metadata_extraction.py:369:5: C901 `_build_stages` is too complex (14 > 10)

Check failure on line 369 in examples/audio/metadata_extraction/run_metadata_extraction.py

View workflow job for this annotation

GitHub Actions / ruff

Ruff (PLR0915)

examples/audio/metadata_extraction/run_metadata_extraction.py:369:5: PLR0915 Too many statements (52 > 50)

Check failure on line 369 in examples/audio/metadata_extraction/run_metadata_extraction.py

View workflow job for this annotation

GitHub Actions / ruff

Ruff (PLR0912)

examples/audio/metadata_extraction/run_metadata_extraction.py:369:5: PLR0912 Too many branches (15 > 12)

Check failure on line 369 in examples/audio/metadata_extraction/run_metadata_extraction.py

View workflow job for this annotation

GitHub Actions / ruff

Ruff (C901)

examples/audio/metadata_extraction/run_metadata_extraction.py:369:5: C901 `_build_stages` is too complex (14 > 10)
corpus_filter = [args.corpus] if args.corpus else None
vad_resources = Resources(cpus=1.0)
if args.vad_backend == "tensorrt":
Expand All @@ -376,6 +382,7 @@
read_concurrency=args.read_concurrency,
resampled_output_dir=args.resampled_output_dir,
resampled_subtype=args.resampled_subtype,
max_audio_duration_sec=args.max_audio_duration_sec,
keep_waveform=not args.resampled_output_dir,
),
]
Expand Down Expand Up @@ -540,7 +547,7 @@
return stages


def main() -> None:

Check failure on line 550 in examples/audio/metadata_extraction/run_metadata_extraction.py

View workflow job for this annotation

GitHub Actions / ruff

Ruff (C901)

examples/audio/metadata_extraction/run_metadata_extraction.py:550:5: C901 `main` is too complex (11 > 10)

Check failure on line 550 in examples/audio/metadata_extraction/run_metadata_extraction.py

View workflow job for this annotation

GitHub Actions / ruff

Ruff (C901)

examples/audio/metadata_extraction/run_metadata_extraction.py:550:5: C901 `main` is too complex (11 > 10)
args = _build_arg_parser().parse_args()

if not args.data_config:
Expand Down
68 changes: 60 additions & 8 deletions nemo_curator/stages/audio/io/nemo_speech_reader.py
Original file line number Diff line number Diff line change
Expand Up @@ -405,6 +405,10 @@ class NeMoSpeechReaderStage(ProcessingStage[FileGroupTask, AudioTask]):
max_io_threads: Maximum number of concurrent I/O threads for
loading audio files in ``process_batch``. Only applies to
single-entry (non-tarred) tasks. Defaults to 8.
max_audio_duration_sec: Maximum source-audio duration to process.
Recordings longer than this are emitted as ``read_error`` audit
rows with ``audio_too_long=True`` rather than being decoded.
Defaults to 12 hours; set to 0 or ``None`` to disable the limit.
resampled_output_dir: If set, write resampled 16 kHz mono WAV files
to this directory. The output filename matches the input stem
with a ``.wav`` extension.
Expand All @@ -419,6 +423,7 @@ class NeMoSpeechReaderStage(ProcessingStage[FileGroupTask, AudioTask]):
# Max shards read in parallel. Caps in-flight waveforms so the object store
# doesn't overflow (without it, Ray launches up to one reader task per CPU).
read_concurrency: int = 2
max_audio_duration_sec: float | None = 12 * 60 * 60
resampled_output_dir: str | None = None
resampled_subtype: str = "FLOAT"
keep_waveform: bool = True
Expand Down Expand Up @@ -689,7 +694,16 @@ def _normalize_lang_fields(entry_data: dict[str, Any]) -> None:
for stale in ("language_pred", "language_pred_source", "language_pred_prob"):
entry_data.pop(stale, None)

def _read_error_task(self, task: FileGroupTask) -> AudioTask:
def _duration_exceeds_limit(self, duration: Any) -> bool: # noqa: ANN401
"""Return whether a known duration exceeds the configured limit."""
if self.max_audio_duration_sec is None or self.max_audio_duration_sec <= 0:
return False
try:
return float(duration) > self.max_audio_duration_sec
except (TypeError, ValueError):
return False

def _read_error_task(self, task: FileGroupTask, *, audio_too_long: bool = False) -> AudioTask:
"""Build a read_error placeholder AudioTask for a source that could not be read.

Emitting a placeholder (rather than dropping the task) is what lets a shard
Expand All @@ -712,6 +726,8 @@ def _read_error_task(self, task: FileGroupTask) -> AudioTask:
"original_file": audio_path,
}
)
if audio_too_long:
entry_data["audio_too_long"] = True
if language and "source_lang" not in entry_data:
entry_data["source_lang"] = language
shard_total = task.reader_config.get("shard_total", 0)
Expand All @@ -728,6 +744,14 @@ def _process_single_entry(self, task: FileGroupTask) -> list[AudioTask]:
audio_path = task.data[0]
hint_sr = entry.get("sampling_rate") or entry.get("sample_rate")

# Prefer actual duration when the source manifest provides it. This
# prevents a multi-hour recording from being decoded into memory just
# to discover that it exceeds the reader's safety limit.
source_duration = entry.get("actual_duration", entry.get("duration", entry.get("proposed_duration")))
if self._duration_exceeds_limit(source_duration):
logger.warning(f"Audio exceeds duration limit, emitting read-error placeholder: {audio_path}")
return [self._read_error_task(task, audio_too_long=True)]

try:
audio, sr, duration = self._load_audio(
audio_path,
Expand All @@ -738,6 +762,10 @@ def _process_single_entry(self, task: FileGroupTask) -> list[AudioTask]:
logger.warning(f"Unreadable audio, emitting read-error placeholder: {audio_path} ({exc})")
return [self._read_error_task(task)]

if self._duration_exceeds_limit(duration):
logger.warning(f"Audio exceeds duration limit after decode, emitting read-error placeholder: {audio_path}")
return [self._read_error_task(task, audio_too_long=True)]

# When the manifest entry describes a sub-segment of the source recording
# (segment-level input, e.g. Granary ASR reading metadata_extraction output),
# emit only that slice so downstream ASR receives a short clip instead of the
Expand Down Expand Up @@ -784,6 +812,31 @@ def _process_single_entry(self, task: FileGroupTask) -> list[AudioTask]:

def _build_cut_entry(self, cut: Any, corpus: str, language: str) -> dict[str, Any]: # noqa: ANN401
"""Decode a single cut and return an entry_data dict (or raise on failure)."""
entry_data = dict(cut.custom) if cut.custom else {}
self._normalize_lang_fields(entry_data)
audio_filepath = ""
if cut.recording and cut.recording.sources:
src = cut.recording.sources[0].source
audio_filepath = src if isinstance(src, str) else cut.id

# CutSet inputs expose their duration before audio loading, so apply
# the same guard without materialising a potentially huge waveform.
if self._duration_exceeds_limit(cut.duration):
entry_data.update(
{
"read_error": True,
"audio_too_long": True,
"duration": cut.duration,
"num_channels": 1,
"corpus": corpus,
"audio_filepath": audio_filepath or cut.id,
"original_file": audio_filepath or cut.id,
}
)
if language and "source_lang" not in entry_data:
entry_data["source_lang"] = language
return entry_data

audio = cut.load_audio().squeeze()
if audio.ndim > 1:
audio = audio.mean(axis=0)
Expand All @@ -797,8 +850,6 @@ def _build_cut_entry(self, cut: Any, corpus: str, language: str) -> dict[str, An
audio = librosa.resample(audio, orig_sr=actual_sr, target_sr=target_sr)

audio = np.asarray(audio, dtype=np.float32)
entry_data = dict(cut.custom) if cut.custom else {}
self._normalize_lang_fields(entry_data)
entry_data.update(
{
"sampling_rate": target_sr,
Expand All @@ -809,11 +860,6 @@ def _build_cut_entry(self, cut: Any, corpus: str, language: str) -> dict[str, An
}
)

audio_filepath = ""
if cut.recording and cut.recording.sources:
src = cut.recording.sources[0].source
audio_filepath = src if isinstance(src, str) else cut.id

if self.resampled_output_dir:
source_name = audio_filepath or cut.id
resampled_path = self._write_resampled_wav(audio, target_sr, source_name)
Expand Down Expand Up @@ -943,6 +989,10 @@ class NeMoSpeechAudioReader(CompositeStage[_EmptyTask, AudioTask]):
max_io_threads: Maximum concurrent threads for loading audio
from S3/object storage. Higher values overlap more network
latency but use more memory. Defaults to 8.
max_audio_duration_sec: Maximum source-audio duration to process.
Longer recordings become ``read_error`` audit rows marked
``audio_too_long``. Defaults to 12 hours; 0 or ``None`` disables
the guard.
resampled_output_dir: If set, write resampled 16 kHz mono WAV files
to this directory. The output filename matches the input stem
with a ``.wav`` extension.
Expand All @@ -959,6 +1009,7 @@ class NeMoSpeechAudioReader(CompositeStage[_EmptyTask, AudioTask]):
output_dir: str | None = None
max_io_threads: int = 8
read_concurrency: int = 2
max_audio_duration_sec: float | None = 12 * 60 * 60
resampled_output_dir: str | None = None
resampled_subtype: str = "FLOAT"
keep_waveform: bool = True
Expand All @@ -979,6 +1030,7 @@ def __post_init__(self) -> None:
NeMoSpeechReaderStage(
max_io_threads=self.max_io_threads,
read_concurrency=self.read_concurrency,
max_audio_duration_sec=self.max_audio_duration_sec,
resampled_output_dir=self.resampled_output_dir,
resampled_subtype=self.resampled_subtype,
keep_waveform=self.keep_waveform,
Expand Down
4 changes: 3 additions & 1 deletion nemo_curator/stages/audio/io/nemo_speech_writer.py
Original file line number Diff line number Diff line change
Expand Up @@ -314,11 +314,13 @@ def process(self, task: AudioTask) -> FileGroupTask: # noqa: C901
if task.data.get("read_error"):
manifest_entry = {
"audio_filepath": "",
"duration": 0.0,
"duration": task.data.get("duration", 0.0),
"sample_rate": self.target_sample_rate,
"sampling_rate": self.target_sample_rate,
"read_error": True,
}
if task.data.get("audio_too_long"):
manifest_entry["audio_too_long"] = True
if original_file:
manifest_entry["original_audio_filepath"] = original_file
for key in ("corpus", "shard_id", "source_lang"):
Expand Down
29 changes: 29 additions & 0 deletions tests/stages/audio/io/test_nemo_speech_writer.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,7 @@ def test_read_error_writes_manifest_and_done(self, tmp_path: Path) -> None:
manifest_path = output_dir / "shard_b.jsonl"
entry = json.loads(manifest_path.read_text(encoding="utf-8").strip())
assert entry["read_error"] is True
assert "audio_too_long" not in entry
assert not (output_dir / "shard_b.jsonl.done").is_file()

task2 = AudioTask(
Expand All @@ -118,6 +119,34 @@ def test_read_error_writes_manifest_and_done(self, tmp_path: Path) -> None:
stage.process(task2)
assert (output_dir / "shard_b.jsonl.done").is_file()

def test_audio_too_long_writes_manifest_with_flag(self, tmp_path: Path) -> None:
output_dir = tmp_path / "out"
stage = NeMoSpeechWriterStage(output_dir=str(output_dir), writer_concurrency=1)
stage.setup()

task = AudioTask(
task_id="clip_long",
dataset_name="test",
data={
"read_error": True,
"audio_too_long": True,
"duration": 54000.0, # 15-hour source file
"original_file": "s3://bucket/audio/too_long.m4a",
"audio_filepath": "s3://bucket/audio/too_long.m4a",
},
_metadata={"_shard_key": "shard_c", "_shard_total": 1},
)

stage.process(task)

entry = json.loads((output_dir / "shard_c.jsonl").read_text(encoding="utf-8").strip())
assert entry["read_error"] is True
assert entry["audio_too_long"] is True
assert entry["duration"] == 54000.0 # actual duration preserved for audit
assert entry["original_audio_filepath"] == "s3://bucket/audio/too_long.m4a"
# A single-entry shard completes even when the only row is audio_too_long.
assert (output_dir / "shard_c.jsonl.done").is_file()

def test_distinct_dirs_same_basename_do_not_collide(self, tmp_path: Path) -> None:
output_dir = tmp_path / "out"
stage = NeMoSpeechWriterStage(output_dir=str(output_dir), writer_concurrency=1)
Expand Down
Loading