From b5b9c00ec46393fff200efb3a2346c71d38e0c4f Mon Sep 17 00:00:00 2001 From: Fan Qian Date: Wed, 5 Aug 2026 19:18:32 +0800 Subject: [PATCH 1/2] Add skip long audio --- .../run_metadata_extraction.py | 14 ++++ .../stages/audio/io/nemo_speech_reader.py | 70 ++++++++++++++++--- .../stages/audio/io/nemo_speech_writer.py | 2 + .../audio/io/test_nemo_speech_writer.py | 2 + 4 files changed, 80 insertions(+), 8 deletions(-) diff --git a/examples/audio/metadata_extraction/run_metadata_extraction.py b/examples/audio/metadata_extraction/run_metadata_extraction.py index 65da83e31e..2b701d23ab 100644 --- a/examples/audio/metadata_extraction/run_metadata_extraction.py +++ b/examples/audio/metadata_extraction/run_metadata_extraction.py @@ -131,6 +131,12 @@ def _build_arg_parser() -> argparse.ArgumentParser: "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( @@ -340,6 +346,12 @@ def _build_arg_parser() -> argparse.ArgumentParser: out = ap.add_argument_group("Output") out.add_argument("--target_sample_rate", type=int, default=16000, help="Output sample rate.") + out.add_argument( + "--save_audio", + action=argparse.BooleanOptionalAction, + default=True, + help="Write segmented opus files alongside the JSONL manifest (default: enabled).", + ) ex = ap.add_argument_group("Executor") ex.add_argument( @@ -376,6 +388,7 @@ def _build_stages(args: argparse.Namespace, language_filter: list[str] | None) - 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, ), ] @@ -535,6 +548,7 @@ def _build_stages(args: argparse.Namespace, language_filter: list[str] | None) - output_dir=args.output_dir, target_sample_rate=args.target_sample_rate, writer_concurrency=args.writer_concurrency, + save_audio=args.save_audio, ) ) return stages diff --git a/nemo_curator/stages/audio/io/nemo_speech_reader.py b/nemo_curator/stages/audio/io/nemo_speech_reader.py index ff20193f33..fa2bcfce58 100644 --- a/nemo_curator/stages/audio/io/nemo_speech_reader.py +++ b/nemo_curator/stages/audio/io/nemo_speech_reader.py @@ -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. @@ -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 @@ -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 @@ -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) @@ -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, @@ -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 @@ -784,6 +812,33 @@ 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, + "sampling_rate": cut.recording.sampling_rate, + "sample_rate": cut.recording.sampling_rate, + "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) @@ -797,8 +852,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, @@ -809,11 +862,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) @@ -943,6 +991,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. @@ -959,6 +1011,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 @@ -979,6 +1032,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, diff --git a/nemo_curator/stages/audio/io/nemo_speech_writer.py b/nemo_curator/stages/audio/io/nemo_speech_writer.py index 7478913597..5121f97891 100644 --- a/nemo_curator/stages/audio/io/nemo_speech_writer.py +++ b/nemo_curator/stages/audio/io/nemo_speech_writer.py @@ -319,6 +319,8 @@ def process(self, task: AudioTask) -> FileGroupTask: # noqa: C901 "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"): diff --git a/tests/stages/audio/io/test_nemo_speech_writer.py b/tests/stages/audio/io/test_nemo_speech_writer.py index 86170cc392..6f4ebc8062 100644 --- a/tests/stages/audio/io/test_nemo_speech_writer.py +++ b/tests/stages/audio/io/test_nemo_speech_writer.py @@ -91,6 +91,7 @@ def test_read_error_writes_manifest_and_done(self, tmp_path: Path) -> None: dataset_name="test", data={ "read_error": True, + "audio_too_long": True, "original_file": "s3://bucket/audio/broken.m4a", "audio_filepath": "s3://bucket/audio/broken.m4a", }, @@ -102,6 +103,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 entry["audio_too_long"] is True assert not (output_dir / "shard_b.jsonl.done").is_file() task2 = AudioTask( From 9d88aa4ef44ce3b410dd4679268b8d2c7c3428c1 Mon Sep 17 00:00:00 2001 From: Fan Qian Date: Wed, 5 Aug 2026 19:47:43 +0800 Subject: [PATCH 2/2] Remove arg --- .../run_metadata_extraction.py | 7 ----- .../stages/audio/io/nemo_speech_reader.py | 2 -- .../stages/audio/io/nemo_speech_writer.py | 2 +- .../audio/io/test_nemo_speech_writer.py | 31 +++++++++++++++++-- 4 files changed, 30 insertions(+), 12 deletions(-) diff --git a/examples/audio/metadata_extraction/run_metadata_extraction.py b/examples/audio/metadata_extraction/run_metadata_extraction.py index 2b701d23ab..4e7e1f9628 100644 --- a/examples/audio/metadata_extraction/run_metadata_extraction.py +++ b/examples/audio/metadata_extraction/run_metadata_extraction.py @@ -346,12 +346,6 @@ def _build_arg_parser() -> argparse.ArgumentParser: out = ap.add_argument_group("Output") out.add_argument("--target_sample_rate", type=int, default=16000, help="Output sample rate.") - out.add_argument( - "--save_audio", - action=argparse.BooleanOptionalAction, - default=True, - help="Write segmented opus files alongside the JSONL manifest (default: enabled).", - ) ex = ap.add_argument_group("Executor") ex.add_argument( @@ -548,7 +542,6 @@ def _build_stages(args: argparse.Namespace, language_filter: list[str] | None) - output_dir=args.output_dir, target_sample_rate=args.target_sample_rate, writer_concurrency=args.writer_concurrency, - save_audio=args.save_audio, ) ) return stages diff --git a/nemo_curator/stages/audio/io/nemo_speech_reader.py b/nemo_curator/stages/audio/io/nemo_speech_reader.py index fa2bcfce58..283d23788d 100644 --- a/nemo_curator/stages/audio/io/nemo_speech_reader.py +++ b/nemo_curator/stages/audio/io/nemo_speech_reader.py @@ -826,8 +826,6 @@ def _build_cut_entry(self, cut: Any, corpus: str, language: str) -> dict[str, An { "read_error": True, "audio_too_long": True, - "sampling_rate": cut.recording.sampling_rate, - "sample_rate": cut.recording.sampling_rate, "duration": cut.duration, "num_channels": 1, "corpus": corpus, diff --git a/nemo_curator/stages/audio/io/nemo_speech_writer.py b/nemo_curator/stages/audio/io/nemo_speech_writer.py index 5121f97891..737f4881fd 100644 --- a/nemo_curator/stages/audio/io/nemo_speech_writer.py +++ b/nemo_curator/stages/audio/io/nemo_speech_writer.py @@ -314,7 +314,7 @@ 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, diff --git a/tests/stages/audio/io/test_nemo_speech_writer.py b/tests/stages/audio/io/test_nemo_speech_writer.py index 6f4ebc8062..a7158ea193 100644 --- a/tests/stages/audio/io/test_nemo_speech_writer.py +++ b/tests/stages/audio/io/test_nemo_speech_writer.py @@ -91,7 +91,6 @@ def test_read_error_writes_manifest_and_done(self, tmp_path: Path) -> None: dataset_name="test", data={ "read_error": True, - "audio_too_long": True, "original_file": "s3://bucket/audio/broken.m4a", "audio_filepath": "s3://bucket/audio/broken.m4a", }, @@ -103,7 +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 entry["audio_too_long"] is True + assert "audio_too_long" not in entry assert not (output_dir / "shard_b.jsonl.done").is_file() task2 = AudioTask( @@ -120,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)