Fix Audio object handling in HF datasets (has .array attr, not a plain dict)

This commit is contained in:
Michael Treadgold
2026-06-18 20:30:45 +12:00
parent 0ded3603ce
commit 040e594499
+14 -3
View File
@@ -165,6 +165,7 @@ def process_hf_dataset(
skip_no_text = 0 skip_no_text = 0
skip_no_audio = 0 skip_no_audio = 0
skip_bad_ids = 0 skip_bad_ids = 0
skip_bad_audio = 0
try: try:
with output_jsonl.open("w", encoding="utf-8") as f: with output_jsonl.open("w", encoding="utf-8") as f:
for i, row in enumerate(ds): 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]})") print(f" Row {i}: no audio/file key (keys={list(row.keys())[:5]})")
continue continue
if isinstance(audio_info, dict): # HF Audio feature returns an object with .array / .sampling_rate / .path
# Audio loaded as array — save to temp file 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") audio_array = audio_info.get("array")
sr = audio_info.get("sampling_rate", 24000) sr = audio_info.get("sampling_rate", 24000)
if audio_array is None: if audio_array is None:
skip_bad_audio += 1
continue continue
audio_path = str(audio_tmp / f"utt_{i:06d}.wav") audio_path = str(audio_tmp / f"utt_{i:06d}.wav")
sf.write(audio_path, audio_array, sr, subtype="PCM_16") sf.write(audio_path, audio_array, sr, subtype="PCM_16")
@@ -196,8 +203,12 @@ def process_hf_dataset(
if audio_dir: if audio_dir:
audio_path = str(audio_dir / Path(audio_info).name) audio_path = str(audio_dir / Path(audio_info).name)
if not Path(audio_path).is_file(): if not Path(audio_path).is_file():
skip_bad_audio += 1
continue continue
else: else:
skip_bad_audio += 1
if skip_bad_audio <= 3:
print(f" Row {i}: unexpected audio type: {type(audio_info).__name__}")
continue continue
try: try:
@@ -237,7 +248,7 @@ def process_hf_dataset(
pass pass
print(f"Wrote {count} rows to {output_jsonl}") 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 return count