summaryrefslogtreecommitdiff
path: root/src/mic_clipper/segmenter.py
blob: aa0e3b22e61b8cd8692cac426dc70d678da5ef82 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
"""Pure streaming voice-clip state machine."""

from __future__ import annotations

from collections import deque
from dataclasses import dataclass
from datetime import datetime, timedelta

import numpy as np


SAMPLE_RATE = 16_000


@dataclass(frozen=True)
class Start:
    started_at: datetime
    samples: np.ndarray


@dataclass(frozen=True)
class Audio:
    samples: np.ndarray


@dataclass(frozen=True)
class End:
    pass


class Segmenter:
    """Emit stream events while retaining only the configured pre-roll in memory."""

    def __init__(
        self,
        *,
        pre_roll_seconds: float,
        silence_seconds: float,
        max_segment_seconds: float = 600,
        sample_rate: int = SAMPLE_RATE,
    ) -> None:
        self.sample_rate = sample_rate
        self.pre_roll_samples = round(pre_roll_seconds * sample_rate)
        self.silence_samples_limit = round(silence_seconds * sample_rate)
        self.max_segment_samples = round(max_segment_seconds * sample_rate)
        self._pre_roll: deque[np.ndarray] = deque()
        self._pre_roll_size = 0
        self._active = False
        self._segment_samples = 0
        self._silence_samples = 0

    @property
    def active(self) -> bool:
        return self._active

    def push(
        self, samples: np.ndarray, is_speech: bool, captured_at: datetime
    ) -> tuple[Start | Audio | End, ...]:
        samples = np.asarray(samples, dtype=np.float32)
        if samples.ndim != 1:
            raise ValueError("audio must be a mono, one-dimensional array")
        if not len(samples):
            return ()

        if not self._active:
            if not is_speech:
                self._append_pre_roll(samples)
                return ()
            pre_roll = self._take_pre_roll()
            started_at = captured_at - timedelta(
                seconds=len(pre_roll) / self.sample_rate
            )
            self._active = True
            self._segment_samples = len(pre_roll) + len(samples)
            self._silence_samples = 0
            self._append_pre_roll(samples)
            return (Start(started_at, np.concatenate((pre_roll, samples))),)

        if self._segment_samples + len(samples) > self.max_segment_samples:
            self._append_pre_roll(samples)
            self._active = True
            self._segment_samples = len(samples)
            self._silence_samples = 0
            return (End(), Start(captured_at, samples))

        self._segment_samples += len(samples)
        if is_speech:
            self._silence_samples = 0
        else:
            self._silence_samples += len(samples)
        self._append_pre_roll(samples)

        events: tuple[Start | Audio | End, ...] = (Audio(samples),)
        if self._silence_samples >= self.silence_samples_limit:
            self._active = False
            self._segment_samples = 0
            self._silence_samples = 0
            events += (End(),)
        return events

    def _append_pre_roll(self, samples: np.ndarray) -> None:
        self._pre_roll.append(samples)
        self._pre_roll_size += len(samples)
        while self._pre_roll_size > self.pre_roll_samples:
            excess = self._pre_roll_size - self.pre_roll_samples
            oldest = self._pre_roll[0]
            if len(oldest) <= excess:
                self._pre_roll.popleft()
                self._pre_roll_size -= len(oldest)
            else:
                self._pre_roll[0] = oldest[excess:]
                self._pre_roll_size -= excess

    def _take_pre_roll(self) -> np.ndarray:
        if not self._pre_roll:
            return np.empty(0, dtype=np.float32)
        samples = np.concatenate(tuple(self._pre_roll))
        self._pre_roll.clear()
        self._pre_roll_size = 0
        return samples