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.
- {simulstream-0.3.0/simulstream.egg-info → simulstream-1.0.0}/PKG-INFO +3 -2
- {simulstream-0.3.0 → simulstream-1.0.0}/pyproject.toml +3 -2
- simulstream-1.0.0/simulstream/VERSION.txt +1 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/client/wav_reader_client.py +5 -1
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/latency/mwersegmenter.py +1 -1
- simulstream-1.0.0/simulstream/server/speech_processors/base_doa.py +230 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/base_streamatt.py +39 -10
- 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.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/seamless_streamatt.py +4 -1
- {simulstream-0.3.0 → simulstream-1.0.0/simulstream.egg-info}/PKG-INFO +3 -2
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream.egg-info/SOURCES.txt +6 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream.egg-info/requires.txt +2 -1
- simulstream-1.0.0/uts/client/test_wav_reader_client.py +52 -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-0.3.0/simulstream/version.txt +0 -1
- simulstream-0.3.0/uts/speech_processors/test_streamatt.py +0 -64
- {simulstream-0.3.0 → simulstream-1.0.0}/LICENSE +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/README.md +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/docs/source/conf.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/setup.cfg +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/__init__.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/client/__init__.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/config.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/inference.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/__init__.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/detokenizers.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/logger.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/readers.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/score_latency.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/score_quality.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/__init__.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/latency/__init__.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/latency/stream_laal.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/__init__.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/comet.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/mwersegmenter.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/quality/sacrebleu.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/stats.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/__init__.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/http_server.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/message_processor.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/__init__.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/base.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/canary_sliding_window_retranslation.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/hf_sliding_window_retranslation.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/incremental_output.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/remote/__init__.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/remote/http_proxy_speech_processor.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/remote/http_speech_processor_server.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/seamless_sliding_window_retranslation.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/simuleval_wrapper.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/sliding_window_retranslation.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/vad_wrapper.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/websocket_server.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream.egg-info/dependency_links.txt +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream.egg-info/entry_points.txt +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/simulstream.egg-info/top_level.txt +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/uts/__init__.py +0 -0
- {simulstream-0.3.0/uts/metrics → simulstream-1.0.0/uts/client}/__init__.py +0 -0
- {simulstream-0.3.0/uts/speech_processors → simulstream-1.0.0/uts/metrics}/__init__.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/uts/metrics/log_reader.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/uts/metrics/test_stream_laal.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/uts/metrics/test_tokenize_no_inplace.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/uts/speech_processors/test_simuleval_wrapper.py +0 -0
- {simulstream-0.3.0 → simulstream-1.0.0}/uts/test_inference.py +0 -0
- {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
|
+
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
|
|
@@ -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)
|
{simulstream-0.3.0 → simulstream-1.0.0}/simulstream/metrics/scorers/latency/mwersegmenter.py
RENAMED
|
@@ -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]
|
|
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)
|
{simulstream-0.3.0 → simulstream-1.0.0}/simulstream/server/speech_processors/base_streamatt.py
RENAMED
|
@@ -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
|
|
177
|
-
frames_to_cut = earliest_attended_idx * self.
|
|
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
|
-
|
|
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(
|
|
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
|
|
289
|
-
#
|
|
290
|
-
if
|
|
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]
|