Files
wehub-resource-sync ba4be087d5
Create PR to main with cherry-pick from release / cherry-pick (push) Failing after 0s
CICD NeMo / pre-flight (push) Failing after 0s
CICD NeMo / configure (push) Has been skipped
Build, validate, and release Neural Modules / pre-flight (push) Failing after 1s
CICD NeMo / code-linting (push) Has been skipped
Build, validate, and release Neural Modules / release (push) Has been skipped
Build, validate, and release Neural Modules / release-summary (push) Has been cancelled
CICD NeMo / cicd-test-container-build (push) Has been cancelled
CICD NeMo / cicd-import-tests (push) Has been cancelled
CICD NeMo / L0_Setup_Test_Data_And_Models (push) Has been cancelled
CICD NeMo / cicd-main-unit-tests (push) Has been cancelled
CICD NeMo / cicd-main-speech (push) Has been cancelled
CICD NeMo / Nemo_CICD_Test (push) Has been cancelled
CICD NeMo / Coverage (e2e) (push) Has been cancelled
CICD NeMo / Coverage (unit-test) (push) Has been cancelled
CodeQL / Analyze (python) (push) Has been cancelled
CICD NeMo / cicd-wait-in-queue (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 13:28:58 +08:00

734 lines
29 KiB
Python

# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import random
import re
import torch
import torch.utils.data
from lhotse import CutSet, MonoCut, Recording, Seconds, SupervisionSegment, compute_num_frames
from lhotse.cut import Cut
from lhotse.dataset.collation import collate_audio, collate_vectors
from lhotse.utils import ifnone
from nemo.collections.common.data.lhotse.text_adapters import Formattable
from nemo.collections.common.tokenizers import TokenizerSpec
from nemo.collections.speechlm2.data.force_align import ForceAligner
from nemo.collections.speechlm2.data.s2s_dataset import _strip_timestamps
from nemo.collections.speechlm2.data.utils import get_pad_id
from nemo.collections.speechlm2.parts.augmentation import AudioAugmenter
from nemo.utils import logging
MCQ_VAL_PROMPT = "Answer the following multiple choice question with an explanation for the answer."
class DuplexSTTDataset(torch.utils.data.Dataset):
"""
A dataset for duplex speech-to-text models.
Unlike DuplexS2SDataset, this dataset does not require target audio and is suitable
for training on standard ASR, AST, and SpeechQA datasets in addition to duplex data.
Args:
tokenizer (TokenizerSpec):
Tokenizer for converting text to token IDs. Must support BOS and EOS tokens.
frame_length (Seconds):
Duration of a single frame in seconds.
source_sample_rate (int):
Sample rate for source audio (e.g., 16000 Hz).
input_roles (list[str], optional):
Speaker roles to treat as inputs. Defaults to ["user"].
output_roles (list[str], optional):
Speaker roles to treat as outputs. Defaults to ["agent"].
aug_by_swap_role (bool, optional):
Whether to augment data by swapping user/agent roles. Defaults to False.
Note: Enabling this requires agent audio to be available in cut.custom['target_audio'].
cfg (dict, optional):
Dataset configuration (e.g., word_align_position).
model_cfg (dict, optional):
Model configuration (e.g., predict_user_text, force_align_user_text).
"""
def __init__(
self,
tokenizer: TokenizerSpec,
frame_length: Seconds,
source_sample_rate: int,
input_roles: list[str] = None,
output_roles: list[str] = None,
aug_by_swap_role: bool = False,
cfg: dict = None,
model_cfg: dict = None,
):
self.tokenizer = tokenizer
self.frame_length = frame_length
self.source_sample_rate = source_sample_rate
self.input_roles = set(ifnone(input_roles, ["user"]))
self.output_roles = set(ifnone(output_roles, ["agent"]))
self.aug_by_swap_role = aug_by_swap_role
self.word_align_position = cfg.get("word_align_position", "left") if cfg is not None else "left"
self.predict_user_text = model_cfg.get("predict_user_text", False) if model_cfg is not None else False
self.force_align_user_text = model_cfg.get("force_align_user_text", False) if model_cfg is not None else None
self.force_align_device = model_cfg.get("force_align_device", "cuda") if model_cfg is not None else "cuda"
self.prepend_word_space = cfg.get("prepend_word_space", True) if cfg is not None else True
self.early_interruption_prob = cfg.get("early_interruption_prob", 0.0) if cfg is not None else 0.0
self.add_mcq_val_prompt = cfg.get("add_mcq_val_prompt", False) if cfg is not None else False
self.cfg = cfg
self.model_cfg = model_cfg
self.force_aligner = None
if self.force_align_user_text:
self.force_aligner = ForceAligner(device=self.force_align_device, frame_length=self.frame_length)
self.audio_augmenter = None
if cfg is not None and (
cfg.get('use_noise_aug', None)
or cfg.get('use_room_ir_aug', None)
or cfg.get('use_mic_ir_aug', None)
or cfg.get('use_codec_aug', None)
):
self.audio_augmenter = AudioAugmenter(sample_rate=source_sample_rate)
assert tokenizer.bos is not None, "BOS support in the tokenizer is required."
assert tokenizer.eos is not None, "EOS support in the tokenizer is required."
def _is_augmentation_task(self, task: str) -> bool:
if self.cfg is not None and self.cfg.get('force_use_noise_augmentation', False):
return True
return task not in ('s2s_duplex_overlap_as_s2s_duplex', 'asr')
def _apply_early_interruption_augmentation(
self,
target_tokens: torch.Tensor,
source_tokens: torch.Tensor,
source_audio: torch.Tensor,
source_audio_lens: torch.Tensor,
batch_idx: int,
) -> None:
"""Simulate early interruption by randomly truncating an agent turn with overlap."""
target_seq = target_tokens[batch_idx]
bos_id = self.tokenizer.bos
eos_id = self.tokenizer.eos
pad_id = get_pad_id(self.tokenizer)
overlap_tokens = self.cfg.get("early_interruption_overlap_tokens", 13) if self.cfg is not None else 13
bos_positions = (target_seq == bos_id).nonzero(as_tuple=True)[0]
eos_positions = (target_seq == eos_id).nonzero(as_tuple=True)[0]
if len(bos_positions) == 0 or len(eos_positions) == 0:
return
turns = []
for bos_pos in bos_positions:
matching_eos = eos_positions[eos_positions > bos_pos]
if len(matching_eos) > 0:
eos_pos = matching_eos[0]
turn_tokens = target_seq[bos_pos + 1 : eos_pos]
non_pad_mask = turn_tokens != pad_id
all_non_pad_positions = (bos_pos + 1 + non_pad_mask.nonzero(as_tuple=True)[0]).tolist()
non_pad_positions = [pos for pos in all_non_pad_positions if (eos_pos - pos) > overlap_tokens]
if len(non_pad_positions) > 0:
turns.append(
{'bos_pos': bos_pos.item(), 'eos_pos': eos_pos.item(), 'non_pad_positions': non_pad_positions}
)
if len(turns) == 0:
return
selected_turn = random.choice(turns)
cutoff_pos = random.choice(selected_turn['non_pad_positions'])
original_eos_pos = selected_turn['eos_pos']
new_eos_pos = min(cutoff_pos + overlap_tokens, original_eos_pos)
frames_to_remove = original_eos_pos - cutoff_pos
if frames_to_remove <= 0:
return
# Update target_tokens: place eos at new_eos_pos, shift tail, pad at end
target_tokens[batch_idx, new_eos_pos] = eos_id
seq_len = target_tokens.shape[1]
cont_start_pos = original_eos_pos + overlap_tokens
tail_length = seq_len - (cont_start_pos + 1)
if tail_length > 0:
target_tokens[batch_idx, new_eos_pos + 1 : new_eos_pos + 1 + tail_length] = target_tokens[
batch_idx, cont_start_pos + 1 : cont_start_pos + 1 + tail_length
].clone()
target_tokens[batch_idx, -frames_to_remove:] = pad_id
# Update source_tokens: shift tail (from cutoff_pos)
src_frames_to_remove = original_eos_pos - cutoff_pos
source_seq_len = source_tokens.shape[1]
source_tail_length = source_seq_len - (original_eos_pos + 1)
if source_tail_length > 0:
source_tokens[batch_idx, cutoff_pos + 1 : cutoff_pos + 1 + source_tail_length] = source_tokens[
batch_idx, original_eos_pos + 1 : original_eos_pos + 1 + source_tail_length
].clone()
source_tokens[batch_idx, -src_frames_to_remove:] = pad_id
# Update source_audio: shift and pad with silence
old_source_len = source_audio_lens[batch_idx].item()
new_bos_source_sample = min(int(cutoff_pos * self.frame_length * self.source_sample_rate), old_source_len)
original_eos_source_sample = min(
int(original_eos_pos * self.frame_length * self.source_sample_rate), old_source_len
)
source_tail_audio_length = old_source_len - original_eos_source_sample
if source_tail_audio_length > 0:
source_audio[batch_idx, new_bos_source_sample : new_bos_source_sample + source_tail_audio_length] = (
source_audio[batch_idx, original_eos_source_sample:old_source_len].clone()
)
source_samples_to_remove = original_eos_source_sample - new_bos_source_sample
if new_bos_source_sample + source_tail_audio_length < source_audio.shape[1]:
source_audio[
batch_idx,
new_bos_source_sample
+ source_tail_audio_length : new_bos_source_sample
+ source_tail_audio_length
+ source_samples_to_remove,
] = 0
def _create_minimal_batch(self) -> dict:
"""Create a minimal valid batch when all cuts are filtered out."""
return {
"sample_id": ["empty_batch"],
"source_audio": torch.zeros((1, 1000), dtype=torch.float32),
"source_audio_lens": torch.tensor([1000], dtype=torch.long),
"target_tokens": torch.full((1, 50), get_pad_id(self.tokenizer), dtype=torch.long),
"target_token_lens": torch.tensor([1], dtype=torch.long),
"source_tokens": torch.full((1, 50), get_pad_id(self.tokenizer), dtype=torch.long),
"source_token_lens": torch.tensor([1], dtype=torch.long),
"source_texts": [""],
"target_texts": [""],
"task": ["s2s_duplex"],
}
def __getitem__(self, all_cuts: CutSet) -> dict:
cuts = all_cuts.filter(lambda c: isinstance(c, Cut))
audio_data = None
if cuts and getattr(cuts[0], 'task', None) == 'asr':
filtered_cuts = []
skipped_cuts = []
for cut in cuts:
if self._has_valid_input(cut):
filtered_cuts.append(cut)
else:
skipped_cuts.append(cut.id)
if skipped_cuts:
logging.info(
f"Skipped {len(skipped_cuts)} cuts with empty input text. Skipped cut ids: {', '.join(skipped_cuts)}"
)
if not filtered_cuts:
logging.warning(
f"All cuts were filtered out! Original batch size: {len(cuts)}. Returning minimal valid batch."
)
return self._create_minimal_batch()
cuts = CutSet.from_cuts(filtered_cuts)
if cuts:
swapped_cuts = []
if self.aug_by_swap_role:
for cut in cuts:
total_turns = cut.custom.get('total_turns', len(cut.supervisions))
if total_turns > 4 and total_turns % 2 == 0:
swapped_cut = self._create_role_swapped_cut(cut)
if swapped_cut:
swapped_cuts.append(swapped_cut)
if swapped_cuts:
all_cuts_combined = CutSet.from_cuts(list(cuts) + swapped_cuts)
else:
all_cuts_combined = cuts
prompt_tokens, prompt_token_lens = collate_system_prompt(
all_cuts_combined, self.tokenizer, add_mcq_val_prompt=self.add_mcq_val_prompt
)
source_audio, source_audio_lens = collate_audio(all_cuts_combined.resample(self.source_sample_rate))
target_tokens, target_token_lens = collate_token_channel(
all_cuts_combined,
self.tokenizer,
self.frame_length,
roles=self.output_roles,
bos_id=self.tokenizer.bos,
eos_id=self.tokenizer.eos,
remove_timestamps=True,
)
# Force align user text (runs in dataloader worker, overlapped with training)
if self.force_align_user_text and torch.is_grad_enabled():
all_cuts_combined = self.force_aligner.batch_force_align_user_audio(
all_cuts_combined, source_sample_rate=self.source_sample_rate
)
source_tokens, source_token_lens = collate_token_channel(
all_cuts_combined,
self.tokenizer,
self.frame_length,
roles=self.input_roles,
bos_id=self.tokenizer.bos,
eos_id=self.tokenizer.eos,
word_align_position=self.word_align_position,
remove_timestamps=not self.predict_user_text,
prepend_word_space=self.prepend_word_space,
)
# Audio augmentation (runs in dataloader workers for performance)
if (
self.audio_augmenter is not None
and torch.is_grad_enabled()
and self._is_augmentation_task(getattr(all_cuts_combined[0], 'task', 's2s_duplex'))
):
source_audio = self.audio_augmenter.augment_batch(self.cfg, source_audio, source_audio_lens)
# Early interruption augmentation
if self.early_interruption_prob > 0 and torch.is_grad_enabled():
for batch_idx in range(target_tokens.shape[0]):
if random.random() < self.early_interruption_prob:
self._apply_early_interruption_augmentation(
target_tokens,
source_tokens,
source_audio,
source_audio_lens,
batch_idx,
)
audio_data = {
"sample_id": [str(cut.id) for cut in all_cuts_combined],
"source_audio": source_audio,
"source_audio_lens": source_audio_lens,
"target_tokens": target_tokens,
"target_token_lens": target_token_lens,
"source_tokens": source_tokens,
"source_token_lens": source_token_lens,
"source_texts": [
" ".join(_strip_timestamps(s.text) for s in cut.supervisions if s.speaker in self.input_roles)
for cut in all_cuts_combined
],
"target_texts": [
" ".join(s.text for s in cut.supervisions if s.speaker in self.output_roles)
for cut in all_cuts_combined
],
"task": [getattr(cut, "task", "s2s_duplex") for cut in all_cuts_combined],
}
if torch.sum(prompt_token_lens) > 0:
audio_data['prompt_tokens'] = prompt_tokens
audio_data['prompt_token_lens'] = prompt_token_lens
text_cuts = all_cuts.filter(lambda c: isinstance(c, Formattable))
text_data = None
if text_cuts:
text_tokens = []
text_token_lens = []
for c in text_cuts:
text_ids = c.input_ids
text_tokens.append(text_ids)
text_token_lens.append(text_ids.shape[0])
text_tokens = collate_vectors(text_tokens, padding_value=get_pad_id(self.tokenizer))
text_token_lens = torch.tensor(text_token_lens, dtype=torch.long)
text_data = {
"text_tokens": text_tokens,
"text_token_lens": text_token_lens,
}
return {
"audio_data": audio_data,
"text_data": text_data,
}
def _create_role_swapped_cut(self, cut):
from io import BytesIO
import numpy as np
import soundfile as sf
from lhotse import AudioSource
assert (
'target_audio' in cut.custom
), f"Role swapping requires target_audio in cut.custom, but cut {cut.id} does not have it. Disable aug_by_swap_role or ensure your data includes target audio."
swapped_supervisions = []
for sup in cut.supervisions:
if sup.speaker == 'User':
new_speaker = 'Assistant'
elif sup.speaker == 'Assistant':
new_speaker = 'User'
else:
continue
swapped_sup = SupervisionSegment(
id=sup.id + "_swapped",
recording_id=sup.recording_id,
start=sup.start,
duration=sup.duration,
channel=sup.channel,
text=sup.text,
language=sup.language,
speaker=new_speaker,
gender=sup.gender,
custom=sup.custom,
alignment=sup.alignment,
)
swapped_supervisions.append(swapped_sup)
swapped_supervisions = sorted(swapped_supervisions, key=lambda s: s.start)
first_agent_idx = None
last_user_idx = None
for i, sup in enumerate(swapped_supervisions):
if sup.speaker == 'Assistant' and first_agent_idx is None:
first_agent_idx = i
if sup.speaker == 'User':
last_user_idx = i
filtered_supervisions = []
for i, sup in enumerate(swapped_supervisions):
if i != first_agent_idx and i != last_user_idx:
filtered_supervisions.append(sup)
if not filtered_supervisions:
return None
first_remaining_start = filtered_supervisions[0].start
adjusted_supervisions = []
for sup in filtered_supervisions:
adjusted_sup = SupervisionSegment(
id=sup.id,
recording_id=sup.recording_id,
start=sup.start - first_remaining_start,
duration=sup.duration,
channel=sup.channel,
text=sup.text,
language=sup.language,
speaker=sup.speaker,
gender=sup.gender,
custom=sup.custom,
alignment=sup.alignment,
)
adjusted_supervisions.append(adjusted_sup)
total_duration = max(s.start + s.duration for s in adjusted_supervisions)
total_samples = int(total_duration * cut.sampling_rate)
new_source_audio = np.zeros(total_samples, dtype=np.float32)
for sup in adjusted_supervisions:
start_sample = int(sup.start * cut.sampling_rate)
end_sample = int((sup.start + sup.duration) * cut.sampling_rate)
if sup.speaker == 'User':
# New User was originally Assistant — audio lives in target_audio channel
original_start = sup.start + first_remaining_start
agent_audio = (
cut.custom['target_audio']
.to_cut()
.truncate(offset=original_start, duration=sup.duration)
.load_audio()
)
if len(agent_audio.shape) > 1:
agent_audio = agent_audio.squeeze()
actual_end = min(end_sample, start_sample + len(agent_audio))
new_source_audio[start_sample:actual_end] = agent_audio[: actual_end - start_sample]
source_buffer = BytesIO()
sf.write(source_buffer, new_source_audio, cut.sampling_rate, format='wav')
source_buffer.seek(0)
new_source_recording = Recording(
id=f"{cut.id}_swapped_source",
sampling_rate=cut.sampling_rate,
num_samples=len(new_source_audio),
duration=total_duration,
sources=[AudioSource(type="memory", channels=[0], source=source_buffer.getvalue())],
)
swapped_cut = MonoCut(
id=f"{cut.id}_swapped",
start=0,
duration=total_duration,
channel=0,
supervisions=adjusted_supervisions,
recording=new_source_recording,
custom={
**cut.custom,
'total_turns': len(adjusted_supervisions),
'role_swapped': True,
},
)
return swapped_cut
def _has_valid_input(self, cut: Cut) -> bool:
return any(s.text.strip() for s in cut.supervisions if s.speaker in self.input_roles)
def collate_token_channel(
cuts: CutSet,
tokenizer: TokenizerSpec,
frame_length: Seconds,
roles: set[str],
bos_id: int = None,
eos_id: int = None,
word_align_position: str = 'left',
remove_timestamps: bool = False,
prepend_word_space: bool = True,
) -> tuple[torch.Tensor, torch.Tensor]:
pad_id = get_pad_id(tokenizer)
tokens = [
_build_token_channel(
c,
tokenizer=tokenizer,
frame_length=frame_length,
roles=roles,
pad_id=pad_id,
bos_id=bos_id,
eos_id=eos_id,
word_align_position=word_align_position,
remove_timestamps=remove_timestamps,
prepend_word_space=prepend_word_space,
)
for c in cuts
]
token_lens = torch.tensor([len(tt) for tt in tokens])
tokens = collate_vectors(tokens, padding_value=pad_id)
return tokens, token_lens
def collate_system_prompt(
cuts: CutSet,
tokenizer: TokenizerSpec,
add_mcq_val_prompt: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Collate system prompts from cuts stored in cut.custom['system_prompt']."""
pad_id = get_pad_id(tokenizer)
tokens = []
for c in cuts:
if c.custom and c.custom.get("system_prompt", None):
prompt_text = c.custom["system_prompt"]
elif add_mcq_val_prompt:
prompt_text = MCQ_VAL_PROMPT
else:
prompt_text = None
if prompt_text:
tokens.append(
torch.as_tensor(
[tokenizer.bos] + tokenizer.text_to_ids(prompt_text) + [tokenizer.eos], dtype=torch.long
)
)
else:
tokens.append(torch.as_tensor([], dtype=torch.long))
token_lens = torch.tensor([len(tt) for tt in tokens])
tokens = collate_vectors(tokens, padding_value=pad_id)
return tokens, token_lens
def _build_token_channel(
cut: Cut,
tokenizer: TokenizerSpec,
frame_length: Seconds,
roles: set[str],
pad_id: int = -1,
bos_id: int = None,
eos_id: int = None,
word_align_position: str = 'left',
remove_timestamps: bool = False,
prepend_word_space: bool = True,
) -> torch.Tensor:
diagnostic = f"Extra info: {cut.id=}"
if getattr(cut, "shard_origin", None) is not None:
diagnostic = f"{diagnostic} {cut.shard_origin=}"
total = compute_num_frames(cut.duration, frame_length, cut.sampling_rate)
tokens = torch.ones(total, dtype=torch.long) * pad_id
for supervision in cut.supervisions:
if supervision.speaker in roles:
pos = compute_num_frames(supervision.start, frame_length, cut.sampling_rate)
if pos >= len(tokens):
logging.warning(
f"Ill-constructed example: the beginning offset of a supervision {pos} is larger than or equal to the example's length {len(tokens)}. {diagnostic}"
)
continue
eospos = compute_num_frames(supervision.end, frame_length, cut.sampling_rate)
available_frames_for_text = eospos - pos
text = supervision.text
text_ids = torch.as_tensor(
[bos_id]
+ _text_to_ids(
text,
tokenizer,
available_frames_for_text=available_frames_for_text,
word_align_position=word_align_position,
remove_timestamps=remove_timestamps,
prepend_word_space=prepend_word_space,
)
)
if available_frames_for_text > 0 and len(text_ids) > available_frames_for_text:
text_ids = text_ids[:available_frames_for_text]
elif available_frames_for_text <= 0:
text_ids = torch.tensor([], dtype=torch.long)
endpos = pos + len(text_ids)
if endpos > len(tokens):
trunc_len = len(tokens) - pos
logging.warning(
f"Truncating training example's text_ids of length {len(text_ids)} by {trunc_len} because {endpos=} > {len(tokens)=}. {diagnostic}"
)
text_ids = text_ids[:trunc_len]
endpos = pos + len(text_ids)
try:
tokens[pos:endpos] = text_ids
except Exception as e:
raise RuntimeError(f"{tokens.shape=} {pos=} {endpos=} {text_ids.shape=} {diagnostic}") from e
if eospos < len(tokens) and eos_id is not None:
tokens[eospos] = eos_id
return tokens
def _text_to_ids(
text: str,
tokenizer: TokenizerSpec,
_TIMESTAMP_PATTERN_STR=r"<\|(\d+)\|>",
available_frames_for_text=None,
word_align_position='left',
remove_timestamps=False,
prepend_word_space=True,
):
if not remove_timestamps and re.compile(_TIMESTAMP_PATTERN_STR).search(text):
text_ids = _text_with_timestamps_to_ids(
text,
tokenizer,
_TIMESTAMP_PATTERN_STR,
available_frames_for_text,
word_align_position,
prepend_word_space=prepend_word_space,
)
else:
_TIMESTAMP_PATTERN = re.compile(_TIMESTAMP_PATTERN_STR)
text = _TIMESTAMP_PATTERN.sub("", text)
text = " ".join(text.strip().split())
text_ids = tokenizer.text_to_ids(text)
return text_ids
def _text_with_timestamps_to_ids(
text: str,
tokenizer: TokenizerSpec,
_TIMESTAMP_PATTERN_STR=r"<\|(\d+)\|>",
available_frames_for_text=None,
word_align_position='left',
prepend_word_space=True,
) -> list[int]:
text_ids, start_times, end_times, word_lens = _extract_text_and_time_tokens(
text,
tokenizer,
_TIMESTAMP_PATTERN_STR,
prepend_word_space=prepend_word_space,
)
text_ids_with_timestamps = _expand_text_with_timestamps_and_word_lengths(
text_ids,
word_lens,
start_times,
end_times,
available_frames_for_text,
frame_rate=0.08,
pad_id=get_pad_id(tokenizer),
word_align_position=word_align_position,
)
return text_ids_with_timestamps
def _extract_text_and_time_tokens(
text, tokenizer: TokenizerSpec, _TIMESTAMP_PATTERN_STR=r"<\|(\d+)\|>", prepend_word_space=True
):
time_tokens = re.findall(_TIMESTAMP_PATTERN_STR, text)
start_time = [int(time_tokens[i]) for i in range(0, len(time_tokens), 2)]
end_time = [int(time_tokens[i]) for i in range(1, len(time_tokens), 2)]
words = re.sub(_TIMESTAMP_PATTERN_STR, '', text).split()
text_ids = []
word_lens = []
for i, word in enumerate(words):
word_with_space = word if i == 0 or not prepend_word_space else ' ' + word
word_ids = tokenizer.text_to_ids(word_with_space)
word_len = len(word_ids)
text_ids.extend(word_ids)
word_lens.append(word_len)
return text_ids, start_time, end_time, word_lens
def _expand_text_with_timestamps_and_word_lengths(
text_ids,
word_lens,
start_time,
end_time,
available_frames_for_text,
frame_rate=0.08,
pad_id=None,
word_align_position='left',
):
def discretize_time(start_token, speech_frame_rate=0.08, timestamp_frame_rate=0.08):
return int(start_token * timestamp_frame_rate / speech_frame_rate)
if pad_id is None:
raise ValueError("pad_id must be provided.")
max_length = available_frames_for_text
text_ids_with_timestamps = [pad_id] * max_length
cur_word_idx = 0
for word_idx, word_len in enumerate(word_lens):
start_idx = discretize_time(start_time[word_idx], speech_frame_rate=frame_rate)
end_idx = discretize_time(end_time[word_idx], speech_frame_rate=frame_rate)
if word_align_position == 'left':
end_idx = min(start_idx + word_len, end_idx)
elif word_align_position == 'right':
start_idx = max(start_idx, end_idx - word_len)
else:
raise ValueError(f"Unknown word_align_position: {word_align_position}")
word_ids = text_ids[cur_word_idx : cur_word_idx + word_len]
for i in range(start_idx, end_idx + 1):
if i - start_idx < len(word_ids) and i < max_length:
token_id = word_ids[i - start_idx]
text_ids_with_timestamps[i] = token_id
cur_word_idx += word_len
return text_ids_with_timestamps