import dataclasses import re from typing import Dict, List, Optional import numpy as np TEXT = "text" TAG = "tag" HESITATION = "--" ANONYMOUS_SPEAKER_PTN = "speaker_{}" ALLOWED_TRANSCRIPT_TAG = (HESITATION, "[laughter]") def _all_or_none_are_none(*items): success = True first = items[0] is None for item in items[1:]: success = success and (first == (item is None)) return success @dataclasses.dataclass(frozen=True) class Token: value: str type: str # must be either "text" or "tag" speaker_id: Optional[str] = None start_s: Optional[float] = None end_s: Optional[float] = None metadata: Optional[Dict] = None def __post_init__(self): if self.type not in (TEXT, TAG): raise ValueError(f"bad token_type `{self.type}` found. Must be `{TEXT}` or `{TAG}`") if not _all_or_none_are_none(self.start_s, self.end_s): raise ValueError( f"start_s and end_s must both be float or none. Found {self.start_s}, {self.end_s}" ) if self.type == TAG and self.value != "--" and (self.value[:1] != "[" or self.value[-1:] != "]"): raise ValueError( f"meta tags must be surrounded by brackets, with the exception of `--`, " f"but {self.value} was found" ) @property def success(self): return self.start_s is not None and self.end_s is not None @property def exclude_in_transcript(self): return self.type == TAG and self.value not in ALLOWED_TRANSCRIPT_TAG def as_dict(self): """Convert to json serializable dict""" dict_rep = dataclasses.asdict(self) # remove None keys clean_dict_rep = {} for k, v in dict_rep.items(): if v is not None: clean_dict_rep[k] = v return clean_dict_rep @classmethod def from_dict(cls, dict_rep): return cls( dict_rep["value"], dict_rep["type"], start_s=dict_rep.get("start_s"), end_s=dict_rep.get("end_s"), speaker_id=dict_rep.get("speaker_id"), metadata=dict_rep.get("metadata"), ) def add_timestamps(self, start_s, end_s): return Token( self.value, self.type, start_s=start_s, end_s=end_s, speaker_id=self.speaker_id, metadata=self.metadata, ) def add_speaker_id(self, speaker_id): return Token( self.value, self.type, start_s=self.start_s, end_s=self.end_s, speaker_id=speaker_id, metadata=self.metadata, ) def add_metadata(self, metadata): return Token( self.value, self.type, start_s=self.start_s, end_s=self.end_s, speaker_id=self.speaker_id, metadata=metadata, ) def remove_timestamps(self): return Token(self.value, self.type, speaker_id=self.speaker_id, metadata=self.metadata) def remove_speaker_id(self): return Token( self.value, self.type, start_s=self.start_s, end_s=self.end_s, metadata=self.metadata, ) def remove_metadata(self): return Token( self.value, self.type, speaker_id=self.speaker_id, start_s=self.start_s, end_s=self.end_s, ) # make alias methods to_dict = as_dict @dataclasses.dataclass(frozen=True) class Tokens: tokens: List[Token] @staticmethod def _verify_monotonic_timestamps(tokens): timestamps = [] for t in tokens: if t.success: timestamps.append(t.start_s) timestamps.append(t.end_s) if len(timestamps) > 0 and np.min(np.diff(timestamps)) < 0: raise ValueError("token timestamps not monotonically increasing") def __post_init__(self): self._verify_monotonic_timestamps(self.tokens) def __len__(self): return len(self.tokens) def __getitem__(self, val): return self.tokens[val] def __repr__(self): repr_text = self.text[:10] if len(self.text) > 10: repr_text += "..." return f"Tokens(text=`{repr_text}`)" @property def success(self): text_tokens = [t for t in self.tokens if t.type == TEXT] return len(text_tokens) == 0 or any([t.success for t in text_tokens]) @property def text(self): return " ".join(t.value for t in self.tokens) @property def plaintext(self): return " ".join(t.value for t in self.tokens if t.type == TEXT) @property def speaker_turns(self): turns = [] cur_speaker = None tmp = [] for token in self.tokens: if token.speaker_id is not None: if token.speaker_id != cur_speaker and len(tmp) > 0: turns.append( { "speaker_id": cur_speaker, "text": " ".join([t.value for t in tmp]), "plaintext": Tokens(tmp).plaintext, "tokens": [t.as_dict() for t in tmp], } ) tmp = [] cur_speaker = token.speaker_id tmp.append(token) if len(tmp) > 0: turns.append( { "speaker_id": cur_speaker, "text": " ".join([t.value for t in tmp]), "plaintext": Tokens(tmp).plaintext, "tokens": [t.as_dict() for t in tmp], } ) return turns @property def n_speakers(self): return len(set([t.speaker_id for t in self.tokens if t.speaker_id is not None])) @property def exclude_in_transcript(self): return any([t.type == TAG and t.value not in ALLOWED_TRANSCRIPT_TAG for t in self.tokens]) def as_dict(self): """Convert to json serializable dict""" return [token.as_dict() for token in self.tokens] @classmethod def from_dict(cls, token_dicts): return cls([Token.from_dict(token_dict) for token_dict in token_dicts]) @classmethod def from_text(cls, text, speaker_id=None): tokens = [] for token_str in re.findall(r"\[.*?\]|[^\s]+", text): if re.match(r"\[.*\]", token_str) or token_str == HESITATION: token_type = TAG else: token_type = TEXT token = Token(token_str, token_type, speaker_id=speaker_id) tokens.append(token) return cls(tokens) @classmethod def from_speaker_turns(cls, turns): tokens = [] for turn in turns: for token_str in re.findall(r"\[.*?\]|[^\s]+", turn["text"]): if re.match(r"\[.*\]", token_str): token = Token(token_str, TAG) tokens.append(token) elif token_str == HESITATION: token = Token(token_str, TAG, speaker_id=turn["speaker_id"]) tokens.append(token) else: token = Token(token_str, TEXT, speaker_id=turn["speaker_id"]) tokens.append(token) return cls(tokens) def anonymize_speakers(self): cur_speaker_n = 0 speaker_id_map = {} anonymized_tokens = [] for t in self.tokens: if t.speaker_id is None: anonymized_tokens.append(t) continue if t.speaker_id not in speaker_id_map: speaker_id_map[t.speaker_id] = cur_speaker_n cur_speaker_n += 1 anonymized_tokens.append(t.add_speaker_id(speaker_id_map[t.speaker_id])) return Tokens(anonymized_tokens) def remove_timestamps(self): return Tokens([token.remove_timestamps() for token in self.tokens]) def remove_speaker_ids(self): return Tokens([token.remove_speaker_id() for token in self.tokens]) def remove_metadatas(self): return Tokens([token.remove_metadata() for token in self.tokens]) # make alias methods to_dict = as_dict