From 040e5944995ebd59e0d0b5d2f1c0e91d45be4198 Mon Sep 17 00:00:00 2001 From: Michael Treadgold Date: Thu, 18 Jun 2026 20:30:45 +1200 Subject: [PATCH] Fix Audio object handling in HF datasets (has .array attr, not a plain dict) --- preprocess_dataset.py | 17 ++++++++++++++--- 1 file changed, 14 insertions(+), 3 deletions(-) diff --git a/preprocess_dataset.py b/preprocess_dataset.py index 92ba2c9..5470633 100644 --- a/preprocess_dataset.py +++ b/preprocess_dataset.py @@ -165,6 +165,7 @@ def process_hf_dataset( skip_no_text = 0 skip_no_audio = 0 skip_bad_ids = 0 + skip_bad_audio = 0 try: with output_jsonl.open("w", encoding="utf-8") as f: for i, row in enumerate(ds): @@ -183,11 +184,17 @@ def process_hf_dataset( print(f" Row {i}: no audio/file key (keys={list(row.keys())[:5]})") continue - if isinstance(audio_info, dict): - # Audio loaded as array — save to temp file + # HF Audio feature returns an object with .array / .sampling_rate / .path + if hasattr(audio_info, "array"): + audio_array = audio_info["array"] if isinstance(audio_info, dict) else audio_info.array + sr = audio_info["sampling_rate"] if isinstance(audio_info, dict) else audio_info.sampling_rate + audio_path = str(audio_tmp / f"utt_{i:06d}.wav") + sf.write(audio_path, audio_array, sr, subtype="PCM_16") + elif isinstance(audio_info, dict): audio_array = audio_info.get("array") sr = audio_info.get("sampling_rate", 24000) if audio_array is None: + skip_bad_audio += 1 continue audio_path = str(audio_tmp / f"utt_{i:06d}.wav") sf.write(audio_path, audio_array, sr, subtype="PCM_16") @@ -196,8 +203,12 @@ def process_hf_dataset( if audio_dir: audio_path = str(audio_dir / Path(audio_info).name) if not Path(audio_path).is_file(): + skip_bad_audio += 1 continue else: + skip_bad_audio += 1 + if skip_bad_audio <= 3: + print(f" Row {i}: unexpected audio type: {type(audio_info).__name__}") continue try: @@ -237,7 +248,7 @@ def process_hf_dataset( pass print(f"Wrote {count} rows to {output_jsonl}") - print(f"Skipped: no_text={skip_no_text} no_audio={skip_no_audio} bad_ids={skip_bad_ids}") + print(f"Skipped: no_text={skip_no_text} no_audio={skip_no_audio} bad_ids={skip_bad_ids} bad_audio={skip_bad_audio}") return count