simulstream 0.3.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 (69) hide show
  1. {simulstream-0.3.0/simulstream.egg-info → simulstream-1.0.0}/PKG-INFO +3 -2
  2. {simulstream-0.3.0 → simulstream-1.0.0}/pyproject.toml +3 -2
  3. simulstream-1.0.0/simulstream/VERSION.txt +1 -0
  4. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/client/wav_reader_client.py +5 -1
  5. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/latency/mwersegmenter.py +1 -1
  6. simulstream-1.0.0/simulstream/server/speech_processors/base_doa.py +230 -0
  7. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/base_streamatt.py +39 -10
  8. simulstream-1.0.0/simulstream/server/speech_processors/canary_streamatt.py +158 -0
  9. simulstream-1.0.0/simulstream/server/speech_processors/phi4multimodal_doa.py +97 -0
  10. simulstream-1.0.0/simulstream/server/speech_processors/qwenomni_doa.py +152 -0
  11. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/seamless_streamatt.py +4 -1
  12. {simulstream-0.3.0 → simulstream-1.0.0/simulstream.egg-info}/PKG-INFO +3 -2
  13. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream.egg-info/SOURCES.txt +6 -0
  14. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream.egg-info/requires.txt +2 -1
  15. simulstream-1.0.0/uts/client/test_wav_reader_client.py +52 -0
  16. simulstream-1.0.0/uts/speech_processors/__init__.py +0 -0
  17. simulstream-1.0.0/uts/speech_processors/test_streamatt.py +166 -0
  18. simulstream-0.3.0/simulstream/version.txt +0 -1
  19. simulstream-0.3.0/uts/speech_processors/test_streamatt.py +0 -64
  20. {simulstream-0.3.0 → simulstream-1.0.0}/LICENSE +0 -0
  21. {simulstream-0.3.0 → simulstream-1.0.0}/README.md +0 -0
  22. {simulstream-0.3.0 → simulstream-1.0.0}/docs/source/conf.py +0 -0
  23. {simulstream-0.3.0 → simulstream-1.0.0}/setup.cfg +0 -0
  24. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/__init__.py +0 -0
  25. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/client/__init__.py +0 -0
  26. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/config.py +0 -0
  27. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/inference.py +0 -0
  28. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/__init__.py +0 -0
  29. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/detokenizers.py +0 -0
  30. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/logger.py +0 -0
  31. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/readers.py +0 -0
  32. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/score_latency.py +0 -0
  33. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/score_quality.py +0 -0
  34. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/__init__.py +0 -0
  35. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/latency/__init__.py +0 -0
  36. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/latency/stream_laal.py +0 -0
  37. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/__init__.py +0 -0
  38. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/comet.py +0 -0
  39. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/mwersegmenter.py +0 -0
  40. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/sacrebleu.py +0 -0
  41. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/stats.py +0 -0
  42. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/__init__.py +0 -0
  43. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/http_server.py +0 -0
  44. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/message_processor.py +0 -0
  45. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/__init__.py +0 -0
  46. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/base.py +0 -0
  47. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/canary_sliding_window_retranslation.py +0 -0
  48. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/hf_sliding_window_retranslation.py +0 -0
  49. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/incremental_output.py +0 -0
  50. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/remote/__init__.py +0 -0
  51. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/remote/http_proxy_speech_processor.py +0 -0
  52. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/remote/http_speech_processor_server.py +0 -0
  53. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/seamless_sliding_window_retranslation.py +0 -0
  54. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/simuleval_wrapper.py +0 -0
  55. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/sliding_window_retranslation.py +0 -0
  56. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/vad_wrapper.py +0 -0
  57. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/websocket_server.py +0 -0
  58. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream.egg-info/dependency_links.txt +0 -0
  59. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream.egg-info/entry_points.txt +0 -0
  60. {simulstream-0.3.0 → simulstream-1.0.0}/simulstream.egg-info/top_level.txt +0 -0
  61. {simulstream-0.3.0 → simulstream-1.0.0}/uts/__init__.py +0 -0
  62. {simulstream-0.3.0/uts/metrics → simulstream-1.0.0/uts/client}/__init__.py +0 -0
  63. {simulstream-0.3.0/uts/speech_processors → simulstream-1.0.0/uts/metrics}/__init__.py +0 -0
  64. {simulstream-0.3.0 → simulstream-1.0.0}/uts/metrics/log_reader.py +0 -0
  65. {simulstream-0.3.0 → simulstream-1.0.0}/uts/metrics/test_stream_laal.py +0 -0
  66. {simulstream-0.3.0 → simulstream-1.0.0}/uts/metrics/test_tokenize_no_inplace.py +0 -0
  67. {simulstream-0.3.0 → simulstream-1.0.0}/uts/speech_processors/test_simuleval_wrapper.py +0 -0
  68. {simulstream-0.3.0 → simulstream-1.0.0}/uts/test_inference.py +0 -0
  69. {simulstream-0.3.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.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
@@ -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)
@@ -123,7 +123,7 @@ class MWERSegmenterBasedLatencyScorer(LatencyScorer):
123
123
  encoded = [" ".join(self.segmenter.encode(p)) for p in pieces]
124
124
  tokenized_text.append(" ### ".join(encoded))
125
125
  else:
126
- tokenized_text.append(" ".join(self.segmenter.encode(text[i].strip())))
126
+ tokenized_text.append(" ".join(self.segmenter.encode(text[i])))
127
127
  return "\n".join(tokenized_text)
128
128
  else:
129
129
  return "\n".join(text)
@@ -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)
@@ -60,6 +60,10 @@ class BaseStreamAtt(BaseSpeechProcessor):
60
60
  context for next predictions.
61
61
  - **audio_subsampling_factor (int)**: Subsampling factor of the model, if any.
62
62
  Defaults to 1.
63
+ - **mel_hop_samples (int)**: Number of raw waveform samples per mel frame.
64
+ Defaults to 160, i.e. 10ms at 16kHz.
65
+ - **use_raw_audio_history (bool)**: Returns whether ``audio_history`` stores raw
66
+ waveform samples rather than processed frames. Defaults to False.
63
67
  - **text_history_max_len (int)**: The maximum length of the textual history after which
64
68
  the current content is cut. Defaults to 128.
65
69
  - **cross_attention_layer (int)**: Layer from which to extract the cross-attention from.
@@ -75,8 +79,14 @@ class BaseStreamAtt(BaseSpeechProcessor):
75
79
  self.config = config
76
80
  text_history_config = self.config.text_history
77
81
  text_history_cls = class_load(text_history_config.type)
82
+ self.bow_prefix = getattr(self.config, "bow_prefix", BOW_PREFIX)
78
83
  self.text_history_method = text_history_cls(text_history_config)
79
84
  self.audio_subsampling_factor = getattr(self.config, "audio_subsampling_factor", 1)
85
+ self.mel_hop_samples = getattr(self.config, "mel_hop_samples", 160)
86
+ self.use_raw_audio_history = getattr(self.config, "use_raw_audio_history", False)
87
+ self.frames_to_audio_history = self.audio_subsampling_factor
88
+ if self.use_raw_audio_history:
89
+ self.frames_to_audio_history *= self.mel_hop_samples
80
90
  self.text_history_max_len = getattr(self.config, "text_history_max_len", 128)
81
91
  self.cross_attn_layer = getattr(self.config, "cross_attention_layer", 3)
82
92
  self.cutoff_frame_num = getattr(self.config, "cutoff_frame_num", 2)
@@ -173,8 +183,8 @@ class BaseStreamAtt(BaseSpeechProcessor):
173
183
  # Only one token: use the unique most attended frame
174
184
  earliest_attended_idx = most_attended_idxs[0]
175
185
 
176
- # Multiply by the subsampling factor to recover the original number of frames
177
- frames_to_cut = earliest_attended_idx * self.audio_subsampling_factor
186
+ # Multiply by the number of frames/samples corresponding to the audio history
187
+ frames_to_cut = earliest_attended_idx * self.frames_to_audio_history
178
188
 
179
189
  # Cut the unattended audio features
180
190
  self.audio_history = self.audio_history[frames_to_cut:]
@@ -182,8 +192,7 @@ class BaseStreamAtt(BaseSpeechProcessor):
182
192
  # Check audio history not exceeding maximum allowed length
183
193
  self._cut_audio_exceeding_maxlen()
184
194
 
185
- @staticmethod
186
- def _strip_incomplete_words(tokens: List[str]) -> List[str]:
195
+ def _strip_incomplete_words(self, tokens: List[str]) -> List[str]:
187
196
  """
188
197
  Remove last incomplete word(s) from the new hypothesis.
189
198
 
@@ -198,7 +207,7 @@ class BaseStreamAtt(BaseSpeechProcessor):
198
207
  num_tokens_incomplete = 0
199
208
  for tok in reversed(tokens):
200
209
  num_tokens_incomplete += 1
201
- if tok.startswith(BOW_PREFIX):
210
+ if tok.startswith(self.bow_prefix):
202
211
  # slice off the trailing incomplete tokens
203
212
  tokens_to_write = tokens[:-num_tokens_incomplete]
204
213
  break
@@ -273,11 +282,10 @@ class FixedWordsTextHistory:
273
282
  """
274
283
  Fixed Words textual history selection method that retains a pre-defined
275
284
  number of words in the history (*history_words*).
276
-
277
- The current implementation supports only SentencePiece.
278
285
  """
279
286
  def __init__(self, config: SimpleNamespace):
280
287
  self.history_words = getattr(config, "history_words", 20)
288
+ self.bow_prefix = getattr(config, "bow_prefix", BOW_PREFIX)
281
289
  self.config = config
282
290
 
283
291
  def select_text_history(self, text_history: List[str]):
@@ -285,9 +293,9 @@ class FixedWordsTextHistory:
285
293
  new_history = []
286
294
  for token in reversed(text_history):
287
295
  new_history.append(token)
288
- # Check if 'BOW_PREFIX' (space in SentencePiece) is contained in the token,
289
- # meaning that we reached the beginning of the word that should be counted
290
- if BOW_PREFIX in token:
296
+ # Check if bow_prefix is contained in the token, meaning that we reached
297
+ # the beginning of the word that should be counted
298
+ if self.bow_prefix in token:
291
299
  words_to_keep -= 1
292
300
  # When all the words to keep are consumed, the accumulation is stopped
293
301
  # and the prefix is returned
@@ -297,6 +305,27 @@ class FixedWordsTextHistory:
297
305
  return new_history[::-1]
298
306
 
299
307
 
308
+ class FixedCharsTextHistory:
309
+ """
310
+ Character-count-based textual history selection method that retains a pre-defined number of
311
+ tokens in the history (*history_chars*).
312
+
313
+ Recommended for character-level languages (e.g., Chinese, Japanese) where word-boundary
314
+ markers (▁) are sparse, making :class:`FixedWordsTextHistory` ineffective: when few tokens
315
+ carry a BOW prefix, the word counter never reaches *history_words*, so the history is never
316
+ trimmed and the audio history grows without bound, causing AlignAtt to cut all new tokens.
317
+
318
+ Args:
319
+ config (SimpleNamespace): Configuration object with an optional attribute:
320
+ - **history_chars (int)**: Number of tokens to retain. Defaults to 20.
321
+ """
322
+ def __init__(self, config: SimpleNamespace):
323
+ self.history_chars = getattr(config, "history_chars", 20)
324
+
325
+ def select_text_history(self, text_history: List[str]) -> List[str]:
326
+ return text_history[-self.history_chars:]
327
+
328
+
300
329
  class PunctuationTextHistory:
301
330
  """
302
331
  Punctuation textual history selection method that retains the sentence
@@ -0,0 +1,158 @@
1
+ # Copyright 2025 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 torch
16
+ import numpy as np
17
+
18
+ from types import SimpleNamespace
19
+ from typing import List, Tuple
20
+
21
+ import copy
22
+
23
+ from simulstream.server.speech_processors import SAMPLE_RATE
24
+ from simulstream.server.speech_processors.base_streamatt import BaseStreamAtt
25
+
26
+ from nemo.collections.asr.models import ASRModel
27
+ from nemo.collections.asr.parts.submodules.multitask_decoding import (
28
+ MultiTaskDecodingConfig,
29
+ )
30
+ from nemo.collections.asr.models.aed_multitask_models import (
31
+ MultiTaskTranscriptionConfig,
32
+ )
33
+
34
+
35
+ class CanaryStreamAtt(BaseStreamAtt):
36
+ """
37
+ StreamAtt policy implementation for NVIDIA's Canary-v2 model.
38
+
39
+ Args:
40
+ config (SimpleNamespace): Configuration object.
41
+ Supported attributes:
42
+ - **audio_history_max_duration (int)**: Maximum audio history in seconds.
43
+ Defaults to ``30``.
44
+ - **num_beams (int)**: Number of beams to use for beam search decoding.
45
+ Defaults to ``5``.
46
+ """
47
+
48
+ def __init__(self, config: SimpleNamespace):
49
+ super().__init__(config)
50
+ self._audio_history_max_duration = getattr(self.config, "audio_history_max_duration", 30)
51
+
52
+ expected_mel_hop_samples = (
53
+ self.model.cfg.preprocessor.window_stride * self.model.cfg.preprocessor.sample_rate
54
+ )
55
+
56
+ assert self.mel_hop_samples == expected_mel_hop_samples, (
57
+ f"mel_hop_samples is set to {self.mel_hop_samples} in the config, but the loaded "
58
+ f"model's preprocessor uses {expected_mel_hop_samples} samples per mel frame"
59
+ )
60
+
61
+ # Build the transcription config, which is reused for every transcribe() call.
62
+ self.transcription_cfg = MultiTaskTranscriptionConfig(
63
+ batch_size=1,
64
+ return_hypotheses=True,
65
+ enable_chunking=False,
66
+ verbose=False,
67
+ )
68
+
69
+ @property
70
+ def audio_max_len(self) -> int:
71
+ """Maximum audio history length in raw waveform samples."""
72
+ return self._audio_history_max_duration * SAMPLE_RATE
73
+
74
+ def set_source_language(self, language: str) -> None:
75
+ self.src_lang = language
76
+
77
+ def set_target_language(self, language: str) -> None:
78
+ self.tgt_lang = language
79
+
80
+ @classmethod
81
+ def load_model(cls, config: SimpleNamespace):
82
+ if not hasattr(cls, "model") or cls.model is None:
83
+ cls.model = ASRModel.from_pretrained(model_name=config.model_name)
84
+
85
+ # Configure decoding strategy
86
+ multitask_decoding = MultiTaskDecodingConfig()
87
+ multitask_decoding.strategy = "beam"
88
+ multitask_decoding.return_xattn_scores = True
89
+ multitask_decoding.beam.beam_size = getattr(config, "num_beams", 5)
90
+ cls.model.change_decoding_strategy(multitask_decoding)
91
+
92
+ cls.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
93
+ assert cls.model.cfg.preprocessor.sample_rate == SAMPLE_RATE
94
+ cls.model.to(cls.device)
95
+
96
+ def _build_transcription_config(self):
97
+ """
98
+ Return a ``MultiTaskTranscriptionConfig`` whose prompt encodes the current source/target
99
+ languages, task, PNC preference, and forced decoder prefix.
100
+ """
101
+
102
+ default_turns = self.model.prompt.get_default_dialog_slots()
103
+ default_slots = copy.deepcopy(default_turns[0]["slots"])
104
+ default_slots["source_lang"] = self.src_lang
105
+ default_slots["target_lang"] = self.tgt_lang
106
+
107
+ turns = [
108
+ {
109
+ "role": "user", "slots": default_slots
110
+ },
111
+ {
112
+ "role": "user_prefix",
113
+ "slots": {
114
+ "prefix": self.model.tokenizer.tokens_to_text(self.text_history)
115
+ },
116
+ },
117
+ ]
118
+
119
+ cfg_copy = copy.deepcopy(self.transcription_cfg)
120
+ cfg_copy.prompt = turns
121
+
122
+ return cfg_copy
123
+
124
+ def _preprocess(self, waveform: np.ndarray) -> np.ndarray:
125
+ """
126
+ Append the incoming waveform chunk to the raw audio history and return it.
127
+
128
+ Returns:
129
+ np.ndarray: Accumulated raw audio history.
130
+ """
131
+ waveform = waveform.astype(np.float32)
132
+ if self.audio_history is None:
133
+ self.audio_history = waveform
134
+ else:
135
+ self.audio_history = np.concatenate(
136
+ [self.audio_history, waveform])
137
+
138
+ return self.audio_history
139
+
140
+ def _generate(self, speech: np.ndarray) -> Tuple[List[str], torch.Tensor]:
141
+ override_config = self._build_transcription_config()
142
+
143
+ with torch.inference_mode():
144
+ output = self.model.transcribe(audio=speech, override_config=override_config)
145
+
146
+ hypothesis = output[0]
147
+
148
+ token_ids = hypothesis.y_sequence.detach().cpu().tolist()
149
+ tokens = self.model.tokenizer.ids_to_tokens(token_ids)
150
+
151
+ xatt_raw = hypothesis.xatt_scores[self.cross_attn_layer]
152
+ xatt = xatt_raw.mean(dim=0).cpu() # we average over heads
153
+ xatt = self.normalize_attn(xatt)
154
+
155
+ return tokens, xatt
156
+
157
+ def tokens_to_string(self, tokens: List[str]) -> str:
158
+ return self.model.tokenizer.tokens_to_text(tokens)
@@ -0,0 +1,97 @@
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
+ from types import SimpleNamespace
16
+ from typing import List, Any, Dict, Tuple
17
+
18
+ import numpy as np
19
+ import torch
20
+ from transformers import AutoModelForCausalLM, AutoProcessor, GenerationConfig
21
+
22
+ from simulstream.server.speech_processors import SAMPLE_RATE
23
+ from simulstream.server.speech_processors.base_doa import DecoderOnlyAttention, get_language_name
24
+
25
+
26
+ class Phi4MultimodalDOA(DecoderOnlyAttention):
27
+ """
28
+ Decoder-Only Attention agent for Phi4-Multimodal.
29
+ """
30
+
31
+ # Phi-4 special tokens
32
+ _USER_START = "<|user|>"
33
+ _AUDIO_TOKEN = "<|audio_1|>"
34
+ _END_TOKEN = "<|end|>"
35
+ _ASST_START = "<|assistant|>"
36
+
37
+ ENCODER_SUBSAMPLING_FACTOR = 8
38
+ HOP_LENGTH = 160 # 10ms at 16kHz
39
+ AUDIO_SPECIAL_TOKEN_ID = 200011 # _AUDIO_SPECIAL_TOKEN_ID in modeling_phi4mm.py
40
+
41
+ def __init__(self, config: SimpleNamespace):
42
+ super().__init__(config)
43
+ self.audio_subsampling_factor = self.ENCODER_SUBSAMPLING_FACTOR * self.HOP_LENGTH
44
+
45
+ @classmethod
46
+ def load_model(cls, config: SimpleNamespace) -> None:
47
+ model_path = getattr(
48
+ config,
49
+ "hf_model_name",
50
+ getattr(config, "model_path", "microsoft/Phi-4-multimodal-instruct"),
51
+ )
52
+
53
+ cls.processor = AutoProcessor.from_pretrained(model_path, trust_remote_code=True)
54
+ cls.model = AutoModelForCausalLM.from_pretrained(
55
+ model_path,
56
+ device_map="cuda",
57
+ torch_dtype="auto",
58
+ trust_remote_code=True,
59
+ _attn_implementation="eager",
60
+ )
61
+ cls.model.eval()
62
+ cls.generation_config = GenerationConfig.from_pretrained(model_path)
63
+
64
+ def build_prompt(self) -> str:
65
+ filled_prompt = self.prompt.replace("{tgt_lang}", get_language_name(self.tgt_lang))
66
+ raw_prefix = self.build_raw_text_prefix()
67
+ prompt = (
68
+ f"{self._USER_START}{self._AUDIO_TOKEN}"
69
+ f"{filled_prompt}{self._END_TOKEN}"
70
+ f"{self._ASST_START}{raw_prefix}"
71
+ )
72
+ return prompt
73
+
74
+ def build_processor_inputs(self, waveform: np.ndarray) -> dict:
75
+ return self.processor(
76
+ text=self.build_prompt(),
77
+ audios=[(waveform, SAMPLE_RATE)],
78
+ return_tensors="pt",
79
+ )
80
+
81
+ def _do_generate(self, inputs: Dict[str, Any]) -> Tuple[List[str], List[torch.Tensor]]:
82
+ input_len = inputs["input_ids"].shape[1]
83
+ output = self.model.generate(
84
+ **inputs,
85
+ max_new_tokens=self.max_new_tokens,
86
+ generation_config=self.generation_config,
87
+ num_logits_to_keep=1,
88
+ output_attentions=True,
89
+ return_dict_in_generate=True,
90
+ do_sample=False,
91
+ )
92
+ new_tokens = self.processor.tokenizer.convert_ids_to_tokens(
93
+ output.sequences[0, input_len:], skip_special_tokens=True)
94
+ return new_tokens, output.attentions
95
+
96
+ def _find_audio_positions(self, input_ids: torch.Tensor) -> torch.Tensor:
97
+ return (input_ids[0] == self.AUDIO_SPECIAL_TOKEN_ID).nonzero(as_tuple=True)[0]