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.
Files changed (68) hide show
  1. {simulstream-0.2.0/simulstream.egg-info → simulstream-1.0.0}/PKG-INFO +21 -6
  2. {simulstream-0.2.0 → simulstream-1.0.0}/README.md +18 -4
  3. {simulstream-0.2.0 → simulstream-1.0.0}/pyproject.toml +3 -2
  4. simulstream-1.0.0/simulstream/VERSION.txt +1 -0
  5. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/client/wav_reader_client.py +5 -1
  6. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/inference.py +1 -5
  7. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/score_quality.py +21 -2
  8. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/latency/mwersegmenter.py +36 -3
  9. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/comet.py +1 -5
  10. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/mwersegmenter.py +41 -2
  11. simulstream-1.0.0/simulstream/server/speech_processors/base_doa.py +230 -0
  12. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/base_streamatt.py +51 -12
  13. simulstream-1.0.0/simulstream/server/speech_processors/canary_streamatt.py +158 -0
  14. simulstream-1.0.0/simulstream/server/speech_processors/phi4multimodal_doa.py +97 -0
  15. simulstream-1.0.0/simulstream/server/speech_processors/qwenomni_doa.py +152 -0
  16. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/seamless_streamatt.py +4 -1
  17. {simulstream-0.2.0 → simulstream-1.0.0/simulstream.egg-info}/PKG-INFO +21 -6
  18. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream.egg-info/SOURCES.txt +11 -1
  19. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream.egg-info/requires.txt +2 -1
  20. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream.egg-info/top_level.txt +0 -3
  21. simulstream-1.0.0/uts/client/test_wav_reader_client.py +52 -0
  22. simulstream-1.0.0/uts/metrics/test_stream_laal.py +91 -0
  23. simulstream-1.0.0/uts/metrics/test_tokenize_no_inplace.py +124 -0
  24. simulstream-1.0.0/uts/speech_processors/__init__.py +0 -0
  25. simulstream-1.0.0/uts/speech_processors/test_streamatt.py +166 -0
  26. simulstream-1.0.0/uts/test_inference.py +93 -0
  27. simulstream-0.2.0/simulstream/version.txt +0 -1
  28. {simulstream-0.2.0 → simulstream-1.0.0}/LICENSE +0 -0
  29. {simulstream-0.2.0 → simulstream-1.0.0}/docs/source/conf.py +0 -0
  30. {simulstream-0.2.0 → simulstream-1.0.0}/setup.cfg +0 -0
  31. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/__init__.py +0 -0
  32. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/client/__init__.py +0 -0
  33. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/config.py +0 -0
  34. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/__init__.py +0 -0
  35. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/detokenizers.py +0 -0
  36. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/logger.py +0 -0
  37. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/readers.py +0 -0
  38. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/score_latency.py +0 -0
  39. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/__init__.py +0 -0
  40. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/latency/__init__.py +0 -0
  41. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/latency/stream_laal.py +0 -0
  42. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/__init__.py +0 -0
  43. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/sacrebleu.py +0 -0
  44. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/metrics/stats.py +0 -0
  45. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/__init__.py +0 -0
  46. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/http_server.py +0 -0
  47. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/message_processor.py +0 -0
  48. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/__init__.py +0 -0
  49. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/base.py +0 -0
  50. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/canary_sliding_window_retranslation.py +0 -0
  51. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/hf_sliding_window_retranslation.py +0 -0
  52. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/incremental_output.py +0 -0
  53. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/remote/__init__.py +0 -0
  54. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/remote/http_proxy_speech_processor.py +0 -0
  55. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/remote/http_speech_processor_server.py +0 -0
  56. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/seamless_sliding_window_retranslation.py +0 -0
  57. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/simuleval_wrapper.py +0 -0
  58. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/sliding_window_retranslation.py +0 -0
  59. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/speech_processors/vad_wrapper.py +0 -0
  60. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream/server/websocket_server.py +0 -0
  61. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream.egg-info/dependency_links.txt +0 -0
  62. {simulstream-0.2.0 → simulstream-1.0.0}/simulstream.egg-info/entry_points.txt +0 -0
  63. {simulstream-0.2.0 → simulstream-1.0.0}/uts/__init__.py +0 -0
  64. {simulstream-0.2.0/uts/metrics → simulstream-1.0.0/uts/client}/__init__.py +0 -0
  65. {simulstream-0.2.0/uts/speech_processors → simulstream-1.0.0/uts/metrics}/__init__.py +0 -0
  66. {simulstream-0.2.0 → simulstream-1.0.0}/uts/metrics/log_reader.py +0 -0
  67. {simulstream-0.2.0 → simulstream-1.0.0}/uts/speech_processors/test_simuleval_wrapper.py +0 -0
  68. {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.2.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]==2.4.0; extra == "canary"
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 REFERENCE_FILE.txt \
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.txt \
424
- --transcripts TRANSCRIPTS_FILE.txt
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 REFERENCE_FILE.txt \
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.txt \
187
- --transcripts TRANSCRIPTS_FILE.txt
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]==2.4.0",
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 = [basedir + '/' + line.strip() for line in f if line.strip()]
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
- i = 0
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
 
@@ -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
- resegmented_hypos = mweralign.align_texts(
110
- "\n".join([sentence_def.content for sentence_def in sample.reference]),
111
- sample.hypothesis.final_text).split("\n")
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
- try:
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")
@@ -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
- resegmented_hypos = mweralign.align_texts(
83
- "\n".join(sample.reference), sample.hypothesis).split("\n")
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)