185 lines
5.5 KiB
Python
185 lines
5.5 KiB
Python
from dataclasses import asdict, dataclass
|
|
from typing import List, Optional
|
|
|
|
from .chunking import TokenEstimator, TokenEstimatorProtocol
|
|
from .schemas import SourceTranscriptSegment, TranscriptSegment
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class NormalizationSummary:
|
|
source_segment_count: int
|
|
normalized_segment_count: int
|
|
merge_count: int
|
|
max_segment_gap: float
|
|
ellipsis_gap: float
|
|
max_segment_duration: float
|
|
max_segment_tokens: int
|
|
|
|
def to_dict(self) -> dict:
|
|
return asdict(self)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class NormalizationResult:
|
|
transcript: List[TranscriptSegment]
|
|
summary: NormalizationSummary
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class _WorkingSegment:
|
|
speaker: str
|
|
start: float
|
|
end: float
|
|
text: str
|
|
categories: Optional[List[str]]
|
|
order: int
|
|
|
|
|
|
def normalize_transcript(
|
|
segments: List[SourceTranscriptSegment],
|
|
max_segment_gap: float,
|
|
ellipsis_gap: float,
|
|
max_segment_duration: float,
|
|
max_segment_tokens: int,
|
|
estimator: Optional[TokenEstimatorProtocol] = None,
|
|
) -> NormalizationResult:
|
|
token_estimator = TokenEstimator() if estimator is None else estimator
|
|
working = [
|
|
_WorkingSegment(
|
|
speaker=segment.speaker,
|
|
start=segment.start,
|
|
end=segment.end,
|
|
text=segment.text,
|
|
categories=None if segment.categories is None else list(segment.categories),
|
|
order=index,
|
|
)
|
|
for index, segment in enumerate(segments)
|
|
]
|
|
working.sort(key=lambda segment: (segment.start, segment.end, segment.order))
|
|
|
|
merge_count = 0
|
|
while True:
|
|
candidate_index = _shortest_mergeable_gap_index(
|
|
working,
|
|
max_segment_gap,
|
|
ellipsis_gap,
|
|
max_segment_duration,
|
|
max_segment_tokens,
|
|
token_estimator,
|
|
)
|
|
if candidate_index is None:
|
|
break
|
|
left = working[candidate_index]
|
|
right = working[candidate_index + 1]
|
|
working[candidate_index : candidate_index + 2] = [_merge_segments(left, right, ellipsis_gap)]
|
|
merge_count += 1
|
|
|
|
normalized = _assign_ids(working)
|
|
return NormalizationResult(
|
|
transcript=normalized,
|
|
summary=NormalizationSummary(
|
|
source_segment_count=len(segments),
|
|
normalized_segment_count=len(normalized),
|
|
merge_count=merge_count,
|
|
max_segment_gap=max_segment_gap,
|
|
ellipsis_gap=ellipsis_gap,
|
|
max_segment_duration=max_segment_duration,
|
|
max_segment_tokens=max_segment_tokens,
|
|
),
|
|
)
|
|
|
|
|
|
def _shortest_mergeable_gap_index(
|
|
segments: List[_WorkingSegment],
|
|
max_segment_gap: float,
|
|
ellipsis_gap: float,
|
|
max_segment_duration: float,
|
|
max_segment_tokens: int,
|
|
estimator: TokenEstimatorProtocol,
|
|
) -> Optional[int]:
|
|
best_index = None
|
|
best_gap = None
|
|
for index in range(len(segments) - 1):
|
|
left = segments[index]
|
|
right = segments[index + 1]
|
|
gap = right.start - left.end
|
|
if not _can_merge(
|
|
left,
|
|
right,
|
|
gap,
|
|
max_segment_gap,
|
|
ellipsis_gap,
|
|
max_segment_duration,
|
|
max_segment_tokens,
|
|
estimator,
|
|
):
|
|
continue
|
|
if best_gap is None or gap < best_gap:
|
|
best_index = index
|
|
best_gap = gap
|
|
return best_index
|
|
|
|
|
|
def _can_merge(
|
|
left: _WorkingSegment,
|
|
right: _WorkingSegment,
|
|
gap: float,
|
|
max_segment_gap: float,
|
|
ellipsis_gap: float,
|
|
max_segment_duration: float,
|
|
max_segment_tokens: int,
|
|
estimator: TokenEstimatorProtocol,
|
|
) -> bool:
|
|
if left.speaker != right.speaker:
|
|
return False
|
|
if gap < 0 or gap > max_segment_gap:
|
|
return False
|
|
if right.end - left.start > max_segment_duration:
|
|
return False
|
|
merged_text = _joined_text(left.text, right.text, gap, ellipsis_gap)
|
|
return _estimate_prompt_tokens(merged_text, estimator) <= max_segment_tokens
|
|
|
|
|
|
def _merge_segments(left: _WorkingSegment, right: _WorkingSegment, ellipsis_gap: float) -> _WorkingSegment:
|
|
gap = right.start - left.end
|
|
return _WorkingSegment(
|
|
speaker=left.speaker,
|
|
start=left.start,
|
|
end=right.end,
|
|
text=_joined_text(left.text, right.text, gap, ellipsis_gap),
|
|
categories=_merged_categories(left.categories, right.categories),
|
|
order=left.order,
|
|
)
|
|
|
|
|
|
def _joined_text(left_text: str, right_text: str, gap: float, ellipsis_gap: float) -> str:
|
|
joiner = " " if gap <= ellipsis_gap else " ... "
|
|
return f"{left_text.rstrip()}{joiner}{right_text.lstrip()}"
|
|
|
|
|
|
def _estimate_prompt_tokens(text: str, estimator: TokenEstimatorProtocol) -> int:
|
|
return estimator.estimate_json([{"id": 1, "original_text": text}])
|
|
|
|
|
|
def _merged_categories(left: Optional[List[str]], right: Optional[List[str]]) -> Optional[List[str]]:
|
|
merged: List[str] = []
|
|
for category in (left or []) + (right or []):
|
|
if category not in merged:
|
|
merged.append(category)
|
|
return merged or None
|
|
|
|
|
|
def _assign_ids(segments: List[_WorkingSegment]) -> List[TranscriptSegment]:
|
|
ordered = sorted(segments, key=lambda segment: (segment.start, segment.end, segment.order))
|
|
return [
|
|
TranscriptSegment(
|
|
id=index + 1,
|
|
speaker=segment.speaker,
|
|
start=segment.start,
|
|
end=segment.end,
|
|
text=segment.text,
|
|
categories=None if segment.categories is None else list(segment.categories),
|
|
)
|
|
for index, segment in enumerate(ordered)
|
|
]
|