Repository navigation
Expand file tree
/
Copy pathgenerate_subtitle.py
More file actions
167 lines (140 loc) · 5.56 KB
/
Copy pathgenerate_subtitle.py
File metadata and controls
167 lines (140 loc) · 5.56 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
"""
FunASR Subtitle Generator
Generate SRT/VTT subtitles from audio/video files.
Usage:
python generate_subtitle.py input.mp4
python generate_subtitle.py input.wav --format vtt
python generate_subtitle.py meeting.mp3 --spk # with speaker labels
"""
import argparse
import sys
import os
import re
from funasr.cli import _sentence_timestamp_words, merge_subtitle_segments
def clean_text(text):
return re.sub(r'<\|[^|]*\|>', '', text or "").strip()
def format_time_srt(ms):
h = ms // 3600000
m = (ms % 3600000) // 60000
s = (ms % 60000) // 1000
ms_rem = ms % 1000
return f"{h:02d}:{m:02d}:{s:02d},{ms_rem:03d}"
def format_time_vtt(ms):
h = ms // 3600000
m = (ms % 3600000) // 60000
s = (ms % 60000) // 1000
ms_rem = ms % 1000
return f"{h:02d}:{m:02d}:{s:02d}.{ms_rem:03d}"
def timestamp_bounds_ms(result):
bounds = []
for key in ("timestamp", "timestamps"):
for ts in result.get(key, []) or []:
if isinstance(ts, dict):
start = ts.get("start_time", ts.get("start"))
end = ts.get("end_time", ts.get("end"))
if start is None or end is None:
continue
start_ms = int(float(start) * 1000)
end_ms = int(float(end) * 1000)
elif isinstance(ts, (list, tuple)) and len(ts) >= 2:
start_ms = int(ts[0])
end_ms = int(ts[1])
else:
continue
if end_ms > start_ms:
bounds.append((start_ms, end_ms))
if not bounds:
return None
return min(start for start, _ in bounds), max(end for _, end in bounds)
def main():
parser = argparse.ArgumentParser(description="Generate subtitles from audio/video using FunASR")
parser.add_argument("input", help="Audio/video file path")
parser.add_argument("-o", "--output", help="Output file (default: input.srt)")
parser.add_argument("--format", choices=["srt", "vtt"], default="srt")
parser.add_argument(
"--segment-mode",
choices=["readable", "sentence"],
default="readable",
help="Cue grouping: readable (default) or raw model sentence boundaries",
)
parser.add_argument("--model", default="iic/SenseVoiceSmall")
parser.add_argument("--device", default="cuda")
parser.add_argument(
"--max-single-segment-time",
type=int,
default=60000,
metavar="MS",
help="Maximum VAD segment length in milliseconds (default: 60000)",
)
parser.add_argument("--spk", action="store_true", help="Include speaker labels")
parser.add_argument("--lang", default="auto")
args = parser.parse_args()
if not os.path.exists(args.input):
print(f"Error: {args.input} not found")
sys.exit(1)
output_path = args.output or f"{os.path.splitext(args.input)[0]}.{args.format}"
print(f"Input: {args.input}")
print(f"Output: {output_path}")
from funasr import AutoModel
kwargs = {"model": args.model, "vad_model": "fsmn-vad", "punc_model": "ct-punc",
"vad_kwargs": {"max_single_segment_time": args.max_single_segment_time},
"device": args.device, "disable_update": True}
if args.spk:
kwargs["spk_model"] = "cam++"
if "Fun-ASR-Nano" in args.model or "Qwen" in args.model:
kwargs["trust_remote_code"] = True
kwargs["hub"] = "hf"
print("Loading model...")
model = AutoModel(**kwargs)
print("Transcribing...")
generate_kwargs = {
"input": args.input,
"batch_size": 1,
"sentence_timestamp": True,
"output_timestamp": True,
"return_time_stamps": True,
}
if args.lang != "auto":
generate_kwargs["language"] = args.lang
result = model.generate(**generate_kwargs)
segments = []
result_item = result[0]
sentence_words = _sentence_timestamp_words(result_item)
for index, seg in enumerate(result_item.get("sentence_info", []) or []):
text = clean_text(seg.get("sentence") or seg.get("text", ""))
start = int(seg.get("start", 0) or 0)
end = int(seg.get("end", 0) or 0)
if text and end > start:
item = {
"start": start,
"end": end,
"text": text,
"spk": seg.get("spk"),
"timestamp": seg.get("timestamp") or seg.get("timestamps"),
}
if sentence_words[index]:
item["words"] = sentence_words[index]
segments.append(item)
if not segments:
text = clean_text(result_item.get("text", ""))
if text:
start, end = timestamp_bounds_ms(result_item) or (0, 0)
segments.append({"start": start, "end": end, "text": text, "spk": None})
if not segments:
print("No speech detected.")
sys.exit(0)
if args.segment_mode == "readable":
segments = merge_subtitle_segments(segments)
fmt = format_time_srt if args.format == "srt" else format_time_vtt
with open(output_path, "w", encoding="utf-8") as f:
if args.format == "vtt":
f.write("WEBVTT\n\n")
for i, seg in enumerate(segments, 1):
text = f"[Speaker {seg['spk']}] {seg['text']}" if args.spk and seg['spk'] is not None else seg['text']
if args.format == "srt":
f.write(f"{i}\n{fmt(seg['start'])} --> {fmt(seg['end'])}\n{text}\n\n")
else:
f.write(f"{fmt(seg['start'])} --> {fmt(seg['end'])}\n{text}\n\n")
print(f"Done! {len(segments)} subtitles → {output_path}")
if __name__ == "__main__":
main()