simulstream 0.2.0__tar.gz → 1.0.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {simulstream-0.2.0/simulstream.egg-info → simulstream-1.0.0}/PKG-INFO +21 -6
- {simulstream-0.2.0 → simulstream-1.0.0}/README.md +18 -4
- {simulstream-0.2.0 → simulstream-1.0.0}/pyproject.toml +3 -2
- simulstream-1.0.0/simulstream/VERSION.txt +1 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/client/wav_reader_client.py +5 -1
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/inference.py +1 -5
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/score_quality.py +21 -2
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/latency/mwersegmenter.py +36 -3
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/comet.py +1 -5
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/mwersegmenter.py +41 -2
- simulstream-1.0.0/simulstream/server/speech_processors/base_doa.py +230 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/base_streamatt.py +51 -12
- simulstream-1.0.0/simulstream/server/speech_processors/canary_streamatt.py +158 -0
- simulstream-1.0.0/simulstream/server/speech_processors/phi4multimodal_doa.py +97 -0
- simulstream-1.0.0/simulstream/server/speech_processors/qwenomni_doa.py +152 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/seamless_streamatt.py +4 -1
- {simulstream-0.2.0 → simulstream-1.0.0/simulstream.egg-info}/PKG-INFO +21 -6
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream.egg-info/SOURCES.txt +11 -1
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream.egg-info/requires.txt +2 -1
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream.egg-info/top_level.txt +0 -3
- simulstream-1.0.0/uts/client/test_wav_reader_client.py +52 -0
- simulstream-1.0.0/uts/metrics/test_stream_laal.py +91 -0
- simulstream-1.0.0/uts/metrics/test_tokenize_no_inplace.py +124 -0
- simulstream-1.0.0/uts/speech_processors/__init__.py +0 -0
- simulstream-1.0.0/uts/speech_processors/test_streamatt.py +166 -0
- simulstream-1.0.0/uts/test_inference.py +93 -0
- simulstream-0.2.0/simulstream/version.txt +0 -1
- {simulstream-0.2.0 → simulstream-1.0.0}/LICENSE +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/docs/source/conf.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/setup.cfg +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/__init__.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/client/__init__.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/config.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/__init__.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/detokenizers.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/logger.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/readers.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/score_latency.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/__init__.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/latency/__init__.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/latency/stream_laal.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/__init__.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/sacrebleu.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/stats.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/__init__.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/http_server.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/message_processor.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/__init__.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/base.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/canary_sliding_window_retranslation.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/hf_sliding_window_retranslation.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/incremental_output.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/remote/__init__.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/remote/http_proxy_speech_processor.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/remote/http_speech_processor_server.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/seamless_sliding_window_retranslation.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/simuleval_wrapper.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/sliding_window_retranslation.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/vad_wrapper.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/websocket_server.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream.egg-info/dependency_links.txt +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/simulstream.egg-info/entry_points.txt +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/uts/__init__.py +0 -0
- {simulstream-0.2.0/uts/metrics → simulstream-1.0.0/uts/client}/__init__.py +0 -0
- {simulstream-0.2.0/uts/speech_processors → simulstream-1.0.0/uts/metrics}/__init__.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/uts/metrics/log_reader.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/uts/speech_processors/test_simuleval_wrapper.py +0 -0
- {simulstream-0.2.0 → simulstream-1.0.0}/uts/utils.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: simulstream
|
|
3
|
-
Version: 0.
|
|
3
|
+
Version: 1.0.0
|
|
4
4
|
Summary: A server to run simultaneous/streaming experiments and demo
|
|
5
5
|
Author-email: Marco Gaido <mgaido@fbk.eu>, FBK HLT-MT <mt@fbk.eu>
|
|
6
6
|
License: Apache License
|
|
@@ -215,6 +215,7 @@ Requires-Dist: pyyaml>6.0
|
|
|
215
215
|
Requires-Dist: websockets
|
|
216
216
|
Requires-Dist: torch
|
|
217
217
|
Requires-Dist: librosa
|
|
218
|
+
Requires-Dist: pycountry
|
|
218
219
|
Provides-Extra: dev
|
|
219
220
|
Requires-Dist: pytest==7.4.0; extra == "dev"
|
|
220
221
|
Requires-Dist: flake8; extra == "dev"
|
|
@@ -226,7 +227,7 @@ Provides-Extra: hf
|
|
|
226
227
|
Requires-Dist: transformers==4.48.1; extra == "hf"
|
|
227
228
|
Provides-Extra: canary
|
|
228
229
|
Requires-Dist: Cython; extra == "canary"
|
|
229
|
-
Requires-Dist: nemo_toolkit[asr]==
|
|
230
|
+
Requires-Dist: nemo_toolkit[asr]==3.0.0; extra == "canary"
|
|
230
231
|
Provides-Extra: vad
|
|
231
232
|
Requires-Dist: silero-vad; extra == "vad"
|
|
232
233
|
Provides-Extra: eval
|
|
@@ -414,14 +415,15 @@ can score your speech processor by running:
|
|
|
414
415
|
simulstream_score_latency --scorer stream_laal \
|
|
415
416
|
--eval-config config/speech_processor.yaml \
|
|
416
417
|
--log-file metrics.jsonl \
|
|
417
|
-
--reference
|
|
418
|
+
--reference REFERENCES_FILE.tgt \
|
|
418
419
|
--audio-definition YAML_AUDIO_REFERENCES_DEFINITION.yaml
|
|
419
420
|
|
|
420
421
|
simulstream_score_quality --scorer comet \
|
|
421
422
|
--eval-config config/speech_processor.yaml \
|
|
422
423
|
--log-file metrics.jsonl \
|
|
423
|
-
--references REFERENCES_FILE.
|
|
424
|
-
--transcripts TRANSCRIPTS_FILE.
|
|
424
|
+
--references REFERENCES_FILE.tgt \
|
|
425
|
+
--transcripts TRANSCRIPTS_FILE.src \
|
|
426
|
+
--audio-definition YAML_AUDIO_REFERENCES_DEFINITION.yaml
|
|
425
427
|
|
|
426
428
|
simulstream_stats --eval-config config/speech_processor.yaml \
|
|
427
429
|
--log-file metrics.jsonl
|
|
@@ -435,7 +437,20 @@ the selected metric (``--scorer``).
|
|
|
435
437
|
|
|
436
438
|
Similarly, ``simulstream_score_quality`` evaluated the quality
|
|
437
439
|
of the generated outputs against one (or more) reference (and transcript, only for metrics
|
|
438
|
-
requiring them) file(s).
|
|
440
|
+
requiring them) file(s). Here, the `YAML_AUDIO_REFERENCES_DEFINITION.yaml` has the same number of entries (sentence definitions
|
|
441
|
+
in terms of wav file origin, offset and duration) as `REFERENCES_FILE.tgt` and `TRANSCRIPTS_FILE.src`.
|
|
442
|
+
|
|
443
|
+
As an alternative, `simulstream_score_quality` can be run without the `--audio-definition` specification, by using a list of
|
|
444
|
+
files as arguments of `--references` and `--transcripts`. In this case, the name of the files (trimmed of the extension)
|
|
445
|
+
**must be the same** of the audio files used (i.e. the names present in `metrics.jsonl`). For instance:
|
|
446
|
+
|
|
447
|
+
```
|
|
448
|
+
simulstream_score_quality --scorer comet \
|
|
449
|
+
--eval-config config/speech_processor.yaml \
|
|
450
|
+
--log-file metrics.jsonl \
|
|
451
|
+
--references AUDIO1.tgt,AUDIO2.tgt,AUDIO3.tgt \
|
|
452
|
+
--transcripts AUDIO1.src,AUDIO2.src,AUDIO3.src
|
|
453
|
+
```
|
|
439
454
|
|
|
440
455
|
Lastly, ``simulstream_stats`` computes statistics like the computational cost and flickering ratio.
|
|
441
456
|
|
|
@@ -177,14 +177,15 @@ can score your speech processor by running:
|
|
|
177
177
|
simulstream_score_latency --scorer stream_laal \
|
|
178
178
|
--eval-config config/speech_processor.yaml \
|
|
179
179
|
--log-file metrics.jsonl \
|
|
180
|
-
--reference
|
|
180
|
+
--reference REFERENCES_FILE.tgt \
|
|
181
181
|
--audio-definition YAML_AUDIO_REFERENCES_DEFINITION.yaml
|
|
182
182
|
|
|
183
183
|
simulstream_score_quality --scorer comet \
|
|
184
184
|
--eval-config config/speech_processor.yaml \
|
|
185
185
|
--log-file metrics.jsonl \
|
|
186
|
-
--references REFERENCES_FILE.
|
|
187
|
-
--transcripts TRANSCRIPTS_FILE.
|
|
186
|
+
--references REFERENCES_FILE.tgt \
|
|
187
|
+
--transcripts TRANSCRIPTS_FILE.src \
|
|
188
|
+
--audio-definition YAML_AUDIO_REFERENCES_DEFINITION.yaml
|
|
188
189
|
|
|
189
190
|
simulstream_stats --eval-config config/speech_processor.yaml \
|
|
190
191
|
--log-file metrics.jsonl
|
|
@@ -198,7 +199,20 @@ the selected metric (``--scorer``).
|
|
|
198
199
|
|
|
199
200
|
Similarly, ``simulstream_score_quality`` evaluated the quality
|
|
200
201
|
of the generated outputs against one (or more) reference (and transcript, only for metrics
|
|
201
|
-
requiring them) file(s).
|
|
202
|
+
requiring them) file(s). Here, the `YAML_AUDIO_REFERENCES_DEFINITION.yaml` has the same number of entries (sentence definitions
|
|
203
|
+
in terms of wav file origin, offset and duration) as `REFERENCES_FILE.tgt` and `TRANSCRIPTS_FILE.src`.
|
|
204
|
+
|
|
205
|
+
As an alternative, `simulstream_score_quality` can be run without the `--audio-definition` specification, by using a list of
|
|
206
|
+
files as arguments of `--references` and `--transcripts`. In this case, the name of the files (trimmed of the extension)
|
|
207
|
+
**must be the same** of the audio files used (i.e. the names present in `metrics.jsonl`). For instance:
|
|
208
|
+
|
|
209
|
+
```
|
|
210
|
+
simulstream_score_quality --scorer comet \
|
|
211
|
+
--eval-config config/speech_processor.yaml \
|
|
212
|
+
--log-file metrics.jsonl \
|
|
213
|
+
--references AUDIO1.tgt,AUDIO2.tgt,AUDIO3.tgt \
|
|
214
|
+
--transcripts AUDIO1.src,AUDIO2.src,AUDIO3.src
|
|
215
|
+
```
|
|
202
216
|
|
|
203
217
|
Lastly, ``simulstream_stats`` computes statistics like the computational cost and flickering ratio.
|
|
204
218
|
|
|
@@ -14,7 +14,8 @@ dependencies = [
|
|
|
14
14
|
"pyyaml>6.0",
|
|
15
15
|
"websockets",
|
|
16
16
|
"torch",
|
|
17
|
-
"librosa"
|
|
17
|
+
"librosa",
|
|
18
|
+
"pycountry"
|
|
18
19
|
]
|
|
19
20
|
dynamic = ["version"]
|
|
20
21
|
|
|
@@ -52,7 +53,7 @@ hf = [
|
|
|
52
53
|
|
|
53
54
|
canary = [
|
|
54
55
|
"Cython",
|
|
55
|
-
"nemo_toolkit[asr]==
|
|
56
|
+
"nemo_toolkit[asr]==3.0.0",
|
|
56
57
|
]
|
|
57
58
|
|
|
58
59
|
vad = [
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
1.0.0
|
|
@@ -159,7 +159,11 @@ def load_wav_file_list(list_file_path: str) -> List[str]:
|
|
|
159
159
|
"""
|
|
160
160
|
basedir = os.path.dirname(list_file_path)
|
|
161
161
|
with open(list_file_path, 'r') as f:
|
|
162
|
-
wav_files = [
|
|
162
|
+
wav_files = [
|
|
163
|
+
line.strip() if os.path.isabs(line.strip())
|
|
164
|
+
else os.path.abspath(os.path.join(basedir, line.strip()))
|
|
165
|
+
for line in f if line.strip()
|
|
166
|
+
]
|
|
163
167
|
if not wav_files:
|
|
164
168
|
LOGGER.error("No valid WAV files found in the list.")
|
|
165
169
|
exit(1)
|
|
@@ -53,14 +53,10 @@ def process_audio(
|
|
|
53
53
|
# one speech chunk is the following
|
|
54
54
|
samples_per_chunk = int(
|
|
55
55
|
sample_rate * message_processor.speech_processor.speech_chunk_size)
|
|
56
|
-
|
|
56
|
+
|
|
57
57
|
for i in range(0, len(data), samples_per_chunk):
|
|
58
58
|
output = message_processor.process_speech(data[i:i + samples_per_chunk].tobytes())
|
|
59
59
|
LOGGER.debug(f"response: {output}")
|
|
60
|
-
# send last part of the audio
|
|
61
|
-
if i < len(data):
|
|
62
|
-
output = message_processor.process_speech(data[i:].tobytes())
|
|
63
|
-
LOGGER.debug(f"response: {output}")
|
|
64
60
|
|
|
65
61
|
|
|
66
62
|
def run_inference(
|
|
@@ -124,6 +124,19 @@ def cli_main():
|
|
|
124
124
|
--log-file metrics.jsonl \\
|
|
125
125
|
--references ref.en \\
|
|
126
126
|
--transcripts src.it \\
|
|
127
|
+
--audio-definition audio_def.yaml \\
|
|
128
|
+
--scorer sacrebleu
|
|
129
|
+
|
|
130
|
+
Otherwise, the script can be invoked without specifying the `--audio-definition`,
|
|
131
|
+
but in this case the name of the refererence and transcript files (trimmed of
|
|
132
|
+
the extension) must be the same of the audio files used (i.e. the names present
|
|
133
|
+
in `metrics.jsonl`), e.g.:
|
|
134
|
+
|
|
135
|
+
$ python -m simulstream.metrics.score_quality \\
|
|
136
|
+
--eval-config config/speech-processor.yaml \\
|
|
137
|
+
--log-file metrics.jsonl \\
|
|
138
|
+
--references 1.en,2.en \\
|
|
139
|
+
--transcripts 1.it,2.it \\
|
|
127
140
|
--scorer sacrebleu
|
|
128
141
|
"""
|
|
129
142
|
LOGGER.info(f"Simulstream version: {simulstream.__version__}")
|
|
@@ -140,17 +153,23 @@ def cli_main():
|
|
|
140
153
|
"specified, this should be a single file containing all the lines of the audios in "
|
|
141
154
|
"the reference, which should be of the same length of the audio definition. "
|
|
142
155
|
"Otherwise, this should be a list of files, where each contains the lines "
|
|
143
|
-
"corresponding to an audio file."
|
|
156
|
+
"corresponding to an audio file. In the case of being a list of files, the file "
|
|
157
|
+
"stem must match a corresponding transcript for an audio file (if applicable "
|
|
158
|
+
"to the quality metric).")
|
|
144
159
|
parser.add_argument(
|
|
145
160
|
"--transcripts", nargs="+", type=str,
|
|
146
161
|
help="Path to the textual files containing reference transcripts. If `--audio-definition` "
|
|
147
162
|
"is specified, this should be a single file containing all the lines of the audios "
|
|
148
163
|
"in the reference, which should be of the same length of the audio definition. "
|
|
149
164
|
"Otherwise, this should be a list of files, where each contains the lines "
|
|
150
|
-
"corresponding to an audio file."
|
|
165
|
+
"corresponding to an audio file. In the case of being a list of files, the file "
|
|
166
|
+
"stem must match a corresponding reference for an audio file.")
|
|
151
167
|
parser.add_argument(
|
|
152
168
|
"--audio-definition", "-a", type=str, default=None,
|
|
153
169
|
help="Path to the yaml file containing the segment-level audio information.")
|
|
170
|
+
parser.add_argument(
|
|
171
|
+
"--latency-unit", choices=["char", "word"], default="word",
|
|
172
|
+
help="Whether to computed stats based on words or characters. Default: word.")
|
|
154
173
|
parser.add_argument("--scorer", choices=QUALITY_SCORER_REGISTRY.keys(), required=True)
|
|
155
174
|
args, _ = parser.parse_known_args()
|
|
156
175
|
|
{simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/latency/mwersegmenter.py
RENAMED
|
@@ -17,6 +17,7 @@ from dataclasses import dataclass
|
|
|
17
17
|
from typing import List
|
|
18
18
|
|
|
19
19
|
from mweralign import mweralign
|
|
20
|
+
from mweralign.segmenter import CJSegmenter
|
|
20
21
|
|
|
21
22
|
from simulstream.metrics.readers import ReferenceSentenceDefinition, OutputWithDelays, text_items
|
|
22
23
|
from simulstream.metrics.scorers.latency import LatencyScorer, LatencyScoringSample, LatencyScores
|
|
@@ -58,6 +59,7 @@ class MWERSegmenterBasedLatencyScorer(LatencyScorer):
|
|
|
58
59
|
def __init__(self, args):
|
|
59
60
|
super().__init__(args)
|
|
60
61
|
self.latency_unit = args.latency_unit
|
|
62
|
+
self.segmenter = CJSegmenter() if args.latency_unit == "char" else None
|
|
61
63
|
|
|
62
64
|
def requires_reference(self) -> bool:
|
|
63
65
|
return True
|
|
@@ -101,19 +103,50 @@ class MWERSegmenterBasedLatencyScorer(LatencyScorer):
|
|
|
101
103
|
f"Index {index} should have reached end of delays ({len(delays)})"
|
|
102
104
|
return segmented_delays
|
|
103
105
|
|
|
106
|
+
def _tokenize(self, text: List[str]) -> List[str]:
|
|
107
|
+
"""
|
|
108
|
+
Tokenize text using the segmenter.
|
|
109
|
+
|
|
110
|
+
Borrowed from
|
|
111
|
+
https://github.com/mjpost/mweralign/blob/d23a5479/mweralign/mweralign.py#L147
|
|
112
|
+
"""
|
|
113
|
+
if self.segmenter is not None:
|
|
114
|
+
tokenized_text = []
|
|
115
|
+
for i in range(len(text)):
|
|
116
|
+
if " ### " in text[i]:
|
|
117
|
+
pieces = text[i].strip().split(" ### ")
|
|
118
|
+
encoded = [" ".join(self.segmenter.encode(p)) for p in pieces]
|
|
119
|
+
tokenized_text.append(" ### ".join(encoded))
|
|
120
|
+
elif "\t" in text[i]:
|
|
121
|
+
pieces = text[i].strip().split("\t")
|
|
122
|
+
# underlying C++ binary still uses ###
|
|
123
|
+
encoded = [" ".join(self.segmenter.encode(p)) for p in pieces]
|
|
124
|
+
tokenized_text.append(" ### ".join(encoded))
|
|
125
|
+
else:
|
|
126
|
+
tokenized_text.append(" ".join(self.segmenter.encode(text[i])))
|
|
127
|
+
return "\n".join(tokenized_text)
|
|
128
|
+
else:
|
|
129
|
+
return "\n".join(text)
|
|
130
|
+
|
|
104
131
|
def score(self, samples: List[LatencyScoringSample]) -> LatencyScores:
|
|
105
132
|
resegmented_samples = []
|
|
106
133
|
for sample in samples:
|
|
107
134
|
assert sample.reference is not None, "Cannot realign hypothesis to missing reference"
|
|
108
135
|
|
|
109
|
-
|
|
110
|
-
|
|
111
|
-
sample.
|
|
136
|
+
hypo = self._tokenize([sample.hypothesis.final_text])
|
|
137
|
+
refs = self._tokenize(
|
|
138
|
+
[sentence_def.content for sentence_def in sample.reference])
|
|
139
|
+
resegmented_hypos = mweralign.align_texts(refs, hypo).split("\n")
|
|
112
140
|
|
|
113
141
|
assert len(resegmented_hypos) == len(sample.reference), \
|
|
114
142
|
f"Reference ({sample.audio_name}) has mismatched number of target " \
|
|
115
143
|
f"({len(sample.reference)}) and resegmented lines ({len(resegmented_hypos)})"
|
|
116
144
|
|
|
145
|
+
if self.segmenter is not None:
|
|
146
|
+
# segmenter.decode will strip() the spaces, but we need them to align with delays
|
|
147
|
+
resegmented_hypos = [
|
|
148
|
+
hypo.replace(" ", "").replace("_", " ") for hypo in resegmented_hypos]
|
|
149
|
+
|
|
117
150
|
ideal_delays_splits = self._split_delays_by_segmented_text(
|
|
118
151
|
sample.hypothesis.ideal_delays,
|
|
119
152
|
resegmented_hypos)
|
|
@@ -13,17 +13,13 @@
|
|
|
13
13
|
# limitations under the License
|
|
14
14
|
|
|
15
15
|
import argparse
|
|
16
|
-
import sys
|
|
17
16
|
from typing import List
|
|
18
17
|
|
|
19
18
|
from simulstream.metrics.scorers.quality import register_quality_scorer
|
|
20
19
|
from simulstream.metrics.scorers.quality.mwersegmenter import MWERSegmenterBasedQualityScorer, \
|
|
21
20
|
ResegmentedQualityScoringSample
|
|
22
21
|
|
|
23
|
-
|
|
24
|
-
from comet import download_model, load_from_checkpoint
|
|
25
|
-
except ImportError:
|
|
26
|
-
sys.exit("Please install comet first with `pip install unbabel-comet`.")
|
|
22
|
+
from comet import download_model, load_from_checkpoint
|
|
27
23
|
|
|
28
24
|
|
|
29
25
|
@register_quality_scorer("comet")
|
{simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/mwersegmenter.py
RENAMED
|
@@ -17,6 +17,7 @@ from dataclasses import dataclass
|
|
|
17
17
|
from typing import List, Optional
|
|
18
18
|
|
|
19
19
|
from mweralign import mweralign
|
|
20
|
+
from mweralign.segmenter import CJSegmenter
|
|
20
21
|
|
|
21
22
|
from simulstream.metrics.scorers.quality import QualityScorer, QualityScoringSample
|
|
22
23
|
|
|
@@ -56,6 +57,11 @@ class MWERSegmenterBasedQualityScorer(QualityScorer):
|
|
|
56
57
|
... # Compute a custom quality score
|
|
57
58
|
... return ...
|
|
58
59
|
"""
|
|
60
|
+
|
|
61
|
+
def __init__(self, args):
|
|
62
|
+
super().__init__(args)
|
|
63
|
+
self.segmenter = CJSegmenter() if args.latency_unit == "char" else None
|
|
64
|
+
|
|
59
65
|
def requires_reference(self) -> bool:
|
|
60
66
|
return True
|
|
61
67
|
|
|
@@ -75,15 +81,48 @@ class MWERSegmenterBasedQualityScorer(QualityScorer):
|
|
|
75
81
|
"""
|
|
76
82
|
...
|
|
77
83
|
|
|
84
|
+
def _tokenize(self, text: List[str]) -> List[str]:
|
|
85
|
+
"""
|
|
86
|
+
Tokenize text using the segmenter.
|
|
87
|
+
|
|
88
|
+
Borrowed from
|
|
89
|
+
https://github.com/mjpost/mweralign/blob/d23a5479/mweralign/mweralign.py#L147
|
|
90
|
+
"""
|
|
91
|
+
if self.segmenter is not None:
|
|
92
|
+
tokenized_text = []
|
|
93
|
+
for i in range(len(text)):
|
|
94
|
+
if " ### " in text[i]:
|
|
95
|
+
pieces = text[i].strip().split(" ### ")
|
|
96
|
+
encoded = [" ".join(self.segmenter.encode(p)) for p in pieces]
|
|
97
|
+
tokenized_text.append(" ### ".join(encoded))
|
|
98
|
+
elif "\t" in text[i]:
|
|
99
|
+
pieces = text[i].strip().split("\t")
|
|
100
|
+
# underlying C++ binary still uses ###
|
|
101
|
+
encoded = [" ".join(self.segmenter.encode(p)) for p in pieces]
|
|
102
|
+
tokenized_text.append(" ### ".join(encoded))
|
|
103
|
+
else:
|
|
104
|
+
tokenized_text.append(" ".join(self.segmenter.encode(text[i].strip())))
|
|
105
|
+
return "\n".join(tokenized_text)
|
|
106
|
+
else:
|
|
107
|
+
return "\n".join(text)
|
|
108
|
+
|
|
78
109
|
def score(self, samples: List[QualityScoringSample]) -> float:
|
|
79
110
|
resegmented_samples = []
|
|
80
111
|
for sample in samples:
|
|
81
112
|
assert sample.reference is not None, "Cannot realign hypothesis to missing reference"
|
|
82
|
-
|
|
83
|
-
|
|
113
|
+
hypo = self._tokenize([sample.hypothesis])
|
|
114
|
+
refs = self._tokenize(sample.reference)
|
|
115
|
+
resegmented_hypos = mweralign.align_texts(refs, hypo).split("\n")
|
|
116
|
+
|
|
84
117
|
assert len(sample.reference) == len(resegmented_hypos), \
|
|
85
118
|
f"Reference ({sample.audio_name}) has mismatched number of target " \
|
|
86
119
|
f"({len(sample.reference)}) and resegmented lines ({len(resegmented_hypos)})"
|
|
120
|
+
|
|
121
|
+
if self.segmenter is not None:
|
|
122
|
+
# segmenter.decode will strip() the spaces, but we need them to align with delays
|
|
123
|
+
resegmented_hypos = [
|
|
124
|
+
hypo.replace(" ", "").replace("_", " ") for hypo in resegmented_hypos]
|
|
125
|
+
|
|
87
126
|
resegmented_samples.append(ResegmentedQualityScoringSample(
|
|
88
127
|
sample.audio_name,
|
|
89
128
|
resegmented_hypos,
|
|
@@ -0,0 +1,230 @@
|
|
|
1
|
+
# Copyright 2026 FBK
|
|
2
|
+
|
|
3
|
+
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
4
|
+
# you may not use this file except in compliance with the License.
|
|
5
|
+
# You may obtain a copy of the License at
|
|
6
|
+
|
|
7
|
+
# http://www.apache.org/licenses/LICENSE-2.0
|
|
8
|
+
|
|
9
|
+
# Unless required by applicable law or agreed to in writing, software
|
|
10
|
+
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
11
|
+
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
12
|
+
# See the License for the specific language governing permissions and
|
|
13
|
+
# limitations under the License
|
|
14
|
+
|
|
15
|
+
import logging
|
|
16
|
+
from abc import abstractmethod
|
|
17
|
+
from types import SimpleNamespace
|
|
18
|
+
from typing import List, Tuple, Any, Dict
|
|
19
|
+
|
|
20
|
+
import numpy as np
|
|
21
|
+
import pycountry
|
|
22
|
+
import torch
|
|
23
|
+
|
|
24
|
+
from simulstream.server.speech_processors import SAMPLE_RATE
|
|
25
|
+
from simulstream.server.speech_processors.base_streamatt import BaseStreamAtt
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
logger = logging.getLogger(__name__)
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def get_language_name(code: str) -> str:
|
|
32
|
+
"""Return the language name for an ISO 639-1 code, falling back to the code itself."""
|
|
33
|
+
lang = pycountry.languages.get(alpha_2=code)
|
|
34
|
+
if lang is not None:
|
|
35
|
+
return lang.name
|
|
36
|
+
else:
|
|
37
|
+
logger.warning(f"Language code '{code}' not found in the language list. Using language "
|
|
38
|
+
f"code directly in the prompt, but this can lead to unexpected behavior.")
|
|
39
|
+
return code
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
class DecoderOnlyAttention(BaseStreamAtt):
|
|
43
|
+
"""
|
|
44
|
+
Generic Decoder-only Attention-based policy for SpeechLLMs.
|
|
45
|
+
|
|
46
|
+
The class handles:
|
|
47
|
+
- Raw-waveform history accumulation.
|
|
48
|
+
- Greedy generation with ``output_attentions=True``.
|
|
49
|
+
- Building the proxy cross-attention matrix from self-attention weights.
|
|
50
|
+
- Applying the StreamAtt-based policy on the proxy cross-attention matrix.
|
|
51
|
+
|
|
52
|
+
The derived class should implement the following methods:
|
|
53
|
+
- **load_model**: Loads the model and processor.
|
|
54
|
+
- **build_prompt**: Builds the text prompt to use with audio inputs.
|
|
55
|
+
- **build_processor_inputs**: Builds model inputs from the rolling audio history.
|
|
56
|
+
- **_do_generate**: Returns newly-generated tokens and self-attention scores.
|
|
57
|
+
- **_find_audio_positions**: Returns the indices of audio tokens.
|
|
58
|
+
|
|
59
|
+
Args:
|
|
60
|
+
config (SimpleNamespace): Configuration object. The following additional attributes are
|
|
61
|
+
expected:
|
|
62
|
+
- **attn_layer (int)**: Layer from which to extract attention scores. Defaults to 0.
|
|
63
|
+
- **attn_head (int)**: Attention head to use. If not set, attention scores are averaged
|
|
64
|
+
over all heads.
|
|
65
|
+
- **average_attn_over_layers (bool)**: Whether to average the selected attention view
|
|
66
|
+
over all decoder layers. Defaults to True.
|
|
67
|
+
- **audio_history_max_duration (int)**: Maximum raw waveform length to keep in the
|
|
68
|
+
rolling history, in seconds. Defaults to 180.
|
|
69
|
+
- **max_new_tokens (int)**: Maximum tokens to generate per chunk. Defaults to 32.
|
|
70
|
+
"""
|
|
71
|
+
|
|
72
|
+
def __init__(self, config: SimpleNamespace):
|
|
73
|
+
super().__init__(config)
|
|
74
|
+
self.cross_attn_layer = getattr(self.config, "attn_layer", 0)
|
|
75
|
+
self.cross_attn_head = getattr(self.config, "attn_head", None)
|
|
76
|
+
self.average_attn_over_layers = getattr(self.config, "average_attn_over_layers", True)
|
|
77
|
+
self.audio_history_max_duration = getattr(self.config, "audio_history_max_duration", 180)
|
|
78
|
+
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
79
|
+
self.max_new_tokens = getattr(self.config, "max_new_tokens", 32)
|
|
80
|
+
self.prompt = getattr(self.config, "prompt", "Translate the audio to {tgt_lang}:")
|
|
81
|
+
logger.debug("Prompt:\n%s", self.prompt)
|
|
82
|
+
|
|
83
|
+
@property
|
|
84
|
+
def audio_max_len(self) -> int:
|
|
85
|
+
"""Maximum raw-waveform samples to keep in the rolling audio history."""
|
|
86
|
+
return self.audio_history_max_duration * SAMPLE_RATE
|
|
87
|
+
|
|
88
|
+
@abstractmethod
|
|
89
|
+
def load_model(self, config: SimpleNamespace) -> None:
|
|
90
|
+
"""
|
|
91
|
+
Load the model and processor from *config* and assign them to ``self.model`` and
|
|
92
|
+
``self.processor``.
|
|
93
|
+
|
|
94
|
+
The model **must** be loaded with ``output_attentions=True`` (or the equivalent flag for
|
|
95
|
+
the architecture) and ``_attn_implementation="eager"``.
|
|
96
|
+
"""
|
|
97
|
+
...
|
|
98
|
+
|
|
99
|
+
@abstractmethod
|
|
100
|
+
def build_prompt(self) -> str:
|
|
101
|
+
"""
|
|
102
|
+
Return the prompt string to be used with audio tokens.
|
|
103
|
+
"""
|
|
104
|
+
...
|
|
105
|
+
|
|
106
|
+
@abstractmethod
|
|
107
|
+
def build_processor_inputs(self, waveform: np.ndarray) -> Any:
|
|
108
|
+
"""
|
|
109
|
+
Build processor inputs from the entire rolling waveform history (float32, 16 kHz).
|
|
110
|
+
"""
|
|
111
|
+
...
|
|
112
|
+
|
|
113
|
+
@abstractmethod
|
|
114
|
+
def _do_generate(self, inputs: Dict[str, Any]) -> Tuple[List[str], List[torch.Tensor]]:
|
|
115
|
+
"""
|
|
116
|
+
Runs the actual generation from the underlying model and returns the generated tokens and
|
|
117
|
+
the corresponding self-attention scores.
|
|
118
|
+
"""
|
|
119
|
+
...
|
|
120
|
+
|
|
121
|
+
@abstractmethod
|
|
122
|
+
def _find_audio_positions(self, input_ids: torch.Tensor) -> torch.Tensor:
|
|
123
|
+
"""Return token positions corresponding to the encoded audio span."""
|
|
124
|
+
...
|
|
125
|
+
|
|
126
|
+
def _generate(self, waveform: np.ndarray) -> Tuple[List[str], torch.Tensor]:
|
|
127
|
+
"""
|
|
128
|
+
Generate tokens from the given inputs together with the self-attention scores.
|
|
129
|
+
|
|
130
|
+
Returns:
|
|
131
|
+
Tuple[List[str], torch.Tensor]:
|
|
132
|
+
List[str]: A list of generated tokens.
|
|
133
|
+
torch.Tensor: Self-attention scores between speech and text with dimension
|
|
134
|
+
(token_len, audio_len).
|
|
135
|
+
"""
|
|
136
|
+
inputs = self.build_processor_inputs(waveform).to(self.device)
|
|
137
|
+
input_ids = inputs["input_ids"] # (1, input_len)
|
|
138
|
+
input_len = input_ids.shape[1]
|
|
139
|
+
|
|
140
|
+
audio_positions = self._find_audio_positions(input_ids)
|
|
141
|
+
audio_len = audio_positions.shape[0]
|
|
142
|
+
|
|
143
|
+
# Run the actual generate on the underlying model
|
|
144
|
+
new_tokens, attentions = self._do_generate(inputs)
|
|
145
|
+
|
|
146
|
+
# Build proxy cross-attention for the hypothesis (prefix + new_tokens)
|
|
147
|
+
prefill_attn = self.average_attn(attentions[0])
|
|
148
|
+
prefix_len = len(self.text_history) if self.text_history else 0
|
|
149
|
+
if prefix_len > 0:
|
|
150
|
+
# Prefix rows come from the prefill pass
|
|
151
|
+
prefix_rows = prefill_attn[input_len - prefix_len:, :][:, audio_positions]
|
|
152
|
+
else:
|
|
153
|
+
prefix_rows = torch.zeros(0, audio_len, device=self.device)
|
|
154
|
+
|
|
155
|
+
if new_tokens:
|
|
156
|
+
# The prefill pass predicts the first generated token, so its last row corresponds to
|
|
157
|
+
# the first generated token's proxy audio-attention
|
|
158
|
+
first_new_row = prefill_attn[-1:, audio_positions]
|
|
159
|
+
# Other tokens' attention is present in each generation step
|
|
160
|
+
new_rows = [
|
|
161
|
+
self.average_attn(step_attn).squeeze(0)[audio_positions]
|
|
162
|
+
for step_attn in attentions[1:]
|
|
163
|
+
]
|
|
164
|
+
subsequent_new_attn = torch.stack(new_rows, dim=0)
|
|
165
|
+
new_attn = torch.cat([first_new_row, subsequent_new_attn], dim=0)
|
|
166
|
+
else:
|
|
167
|
+
new_attn = torch.zeros(0, audio_len, device=self.device)
|
|
168
|
+
|
|
169
|
+
cross_attn = torch.cat([prefix_rows, new_attn], dim=0)
|
|
170
|
+
cross_attn = self.normalize_attn(cross_attn)
|
|
171
|
+
return new_tokens, cross_attn
|
|
172
|
+
|
|
173
|
+
def set_target_language(self, language: str) -> None:
|
|
174
|
+
self.tgt_lang = language
|
|
175
|
+
|
|
176
|
+
def set_source_language(self, language: str) -> None:
|
|
177
|
+
self.src_lang = language
|
|
178
|
+
|
|
179
|
+
def build_raw_text_prefix(self) -> str:
|
|
180
|
+
return "".join(self.text_history) if self.text_history else ""
|
|
181
|
+
|
|
182
|
+
def _select_attn_from_layer(self, layer_attn: torch.Tensor) -> torch.Tensor:
|
|
183
|
+
# Generation runs one stream at a time, so remove the singleton batch dimension.
|
|
184
|
+
layer_attn = layer_attn.squeeze(0)
|
|
185
|
+
if self.cross_attn_head is None:
|
|
186
|
+
# Default behavior: average over all heads for this layer.
|
|
187
|
+
return layer_attn.mean(dim=0)
|
|
188
|
+
|
|
189
|
+
num_heads = layer_attn.shape[0]
|
|
190
|
+
if self.cross_attn_head < 0 or self.cross_attn_head >= num_heads:
|
|
191
|
+
raise ValueError(
|
|
192
|
+
f"Invalid attn_head={self.cross_attn_head}. Layer has {num_heads} heads."
|
|
193
|
+
)
|
|
194
|
+
return layer_attn[self.cross_attn_head]
|
|
195
|
+
|
|
196
|
+
def average_attn(self, attn) -> torch.Tensor:
|
|
197
|
+
"""
|
|
198
|
+
Average or select attentions according to ``attn_layer``, ``attn_head``, and
|
|
199
|
+
``average_attn_over_layers``.
|
|
200
|
+
|
|
201
|
+
If ``attn_head`` is not set, attention is averaged over heads. If
|
|
202
|
+
``average_attn_over_layers`` is set, the selected per-layer attention view is also
|
|
203
|
+
averaged across layers; otherwise only ``attn_layer`` is used.
|
|
204
|
+
"""
|
|
205
|
+
if self.average_attn_over_layers:
|
|
206
|
+
# Average the per-layer attention view selected by _select_attn_from_layer.
|
|
207
|
+
return torch.stack(
|
|
208
|
+
[self._select_attn_from_layer(layer_attn) for layer_attn in attn],
|
|
209
|
+
dim=0,
|
|
210
|
+
).mean(dim=0)
|
|
211
|
+
return self._select_attn_from_layer(attn[self.cross_attn_layer])
|
|
212
|
+
|
|
213
|
+
def _preprocess(self, waveform: np.float32) -> np.ndarray:
|
|
214
|
+
"""
|
|
215
|
+
Append *waveform* to ``self.audio_history`` and enforce the maximum length.
|
|
216
|
+
"""
|
|
217
|
+
if self.audio_history is None:
|
|
218
|
+
self.audio_history = waveform
|
|
219
|
+
else:
|
|
220
|
+
self.audio_history = np.concatenate([self.audio_history, waveform])
|
|
221
|
+
|
|
222
|
+
if len(self.audio_history) > self.audio_max_len:
|
|
223
|
+
logger.warning("Audio history exceeded %d samples; trimming.", self.audio_max_len)
|
|
224
|
+
self.audio_history = self.audio_history[-self.audio_max_len:]
|
|
225
|
+
|
|
226
|
+
return self.audio_history
|
|
227
|
+
|
|
228
|
+
def tokens_to_string(self, tokens: List[str]) -> str:
|
|
229
|
+
"""Convert a list of decoded tokens to a plain output string."""
|
|
230
|
+
return "".join(tokens)
|