align.py 2.9 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394959697
  1. #!/usr/bin/env python3
  2. """Whisper-based forced alignment for TTS-generated audio.
  3. Uses faster-whisper (preferred) or openai-whisper to produce word-level
  4. timestamps. Outputs JSON to stdout.
  5. Usage:
  6. python3 align.py <audio_path> [--text TEXT] [--language LANG] [--model MODEL]
  7. Output format:
  8. {"ok": true, "words": [{"word": "...", "start": 0.0, "end": 0.5}, ...]}
  9. {"ok": false, "error": "error message"}
  10. """
  11. import sys
  12. import json
  13. import argparse
  14. def transcribe_with_faster_whisper(audio_path, language=None, model_size="base"):
  15. """Transcribe using faster-whisper and return word-level timestamps."""
  16. from faster_whisper import WhisperModel
  17. model = WhisperModel(model_size, device="cpu", compute_type="int8")
  18. segments, _info = model.transcribe(
  19. audio_path,
  20. word_timestamps=True,
  21. language=language,
  22. vad_filter=True,
  23. )
  24. words = []
  25. for segment in segments:
  26. for word_info in segment.words:
  27. words.append({
  28. "word": word_info.word.strip(),
  29. "start": round(word_info.start, 3),
  30. "end": round(word_info.end, 3),
  31. })
  32. return words
  33. def transcribe_with_whisper(audio_path, language=None, model_size="base"):
  34. """Fallback: transcribe using openai-whisper and return word-level timestamps."""
  35. import whisper
  36. model = whisper.load_model(model_size)
  37. result = model.transcribe(
  38. audio_path,
  39. word_timestamps=True,
  40. language=language,
  41. )
  42. words = []
  43. for segment in result.get("segments", []):
  44. for word_info in segment.get("words", []):
  45. words.append({
  46. "word": word_info["word"].strip(),
  47. "start": round(word_info["start"], 3),
  48. "end": round(word_info["end"], 3),
  49. })
  50. return words
  51. def main():
  52. parser = argparse.ArgumentParser(description="Whisper forced alignment")
  53. parser.add_argument("audio_path", help="Path to audio file")
  54. parser.add_argument("--text", default=None, help="Reference text (unused, reserved)")
  55. parser.add_argument("--language", default=None, help="Language hint (e.g. 'zh', 'en')")
  56. parser.add_argument("--model", default="base", help="Whisper model size (tiny/base/small/medium/large)")
  57. args = parser.parse_args()
  58. try:
  59. try:
  60. words = transcribe_with_faster_whisper(
  61. args.audio_path,
  62. language=args.language,
  63. model_size=args.model,
  64. )
  65. except ImportError:
  66. words = transcribe_with_whisper(
  67. args.audio_path,
  68. language=args.language,
  69. model_size=args.model,
  70. )
  71. print(json.dumps({"ok": True, "words": words}))
  72. except Exception as e:
  73. print(json.dumps({"ok": False, "error": str(e)}))
  74. sys.exit(0)
  75. if __name__ == "__main__":
  76. main()