File size: 7,369 Bytes
21fd722
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
"""Bounded PCM reader, explicit source clocks, and append-only byte tokenizer."""
from __future__ import annotations
import codecs
from dataclasses import dataclass
import hashlib
import struct
from pathlib import Path
import numpy as np


def require(ok, message):
    if not ok:
        raise ValueError(message)


@dataclass(frozen=True)
class StreamProfile:
    startup: str = 'realtime18'
    delay_tokens: int = 3
    language_id: int = 151668
    gear: int = 4

    def __post_init__(self):
        require(self.startup in ('hf', 'realtime18'), 'startup must be hf or realtime18')
        require(type(self.delay_tokens) is int and 1 <= self.delay_tokens <= 30, 'delay tokens 1..30')
        require(self.gear == 4, 'initial MLX streaming route supports trained 80 ms gear4 only')
        require(self.language_id in (151667, 151668), 'language token')

    @property
    def period(self): return 320 * self.gear

    @property
    def left_samples(self): return 18 * self.period

    @property
    def prefill(self): return 18 if self.startup == 'realtime18' else 19 + self.delay_tokens

    @property
    def eof_policy(self):
        return 'official_complete_centered_window_delay_plus_11_silence_tokens' if self.startup == 'realtime18' else 'hf_diagnostic_ceil_audio_plus_11_emissions'

    def prompt(self):
        return [151644, self.language_id] + [151665] * (self.prefill - 2)

    def bounds(self, index):
        require(type(index) is int and index >= 0, 'negative/noninteger output index')
        start = 0 if index == 0 else (self.prefill + index - 1) * self.period
        end = (self.prefill + index) * self.period
        return max(0, start - 840), end + 40

    def source_needed(self, index):
        return max(0, self.bounds(index)[1] - self.left_samples)

    def emissions_at_eof(self, samples):
        require(type(samples) is int and samples >= 0, 'invalid source sample count')
        if self.startup == 'realtime18':
            # Official buffer appends (delay+1+10)*P zeros and only emits a
            # complete centered window. Do not round the source to a token.
            return max(0, (samples + (self.delay_tokens + 11) * self.period - 40) // self.period + 1)
        return (samples + self.period - 1) // self.period + 11


class PCMReader:
    """Blocking sequential finite F32LE input; no total-waveform accumulation."""
    def __init__(self, stream):
        self.stream = stream
        self.samples = 0
        self.base = 0
        self.buffer = np.empty(0, dtype=np.float32)
        self.eof = False
        self.digest = hashlib.sha256()
        self.max_retained_samples = 0

    def ensure(self, target):
        require(type(target) is int and target >= 0, 'sample target')
        while self.samples < target and not self.eof:
            desired = min(8192, (target - self.samples) * 4)
            data = bytearray()
            while len(data) < desired:
                block = self.stream.read(desired - len(data))
                if block is None:
                    raise ValueError('nonblocking streams require an external blocking adapter')
                if not block:
                    self.eof = True
                    break
                data.extend(block)
            require(len(data) % 4 == 0, 'truncated F32LE sample at EOF')
            values = np.frombuffer(data, dtype='<f4')
            require(np.isfinite(values).all(), 'nonfinite PCM sample')
            self.digest.update(data)
            self.buffer = np.concatenate([self.buffer, values])
            self.samples += len(values)
            self.max_retained_samples = max(self.max_retained_samples, len(self.buffer))
            require(len(self.buffer) <= 65536, 'PCM retention invariant exceeded')

    def window(self, profile, index):
        start, end = profile.bounds(index)
        source_start = max(0, start - profile.left_samples)
        source_end = max(0, end - profile.left_samples)
        require(source_start >= self.base, 'requested discarded source samples')
        discard = min(source_start - self.base, len(self.buffer))
        self.buffer = self.buffer[discard:].copy()
        self.base += discard
        # After EOF, silence positions can advance past the retained source.
        self.ensure(source_end)
        if self.eof and index >= profile.emissions_at_eof(self.samples):
            return None
        output = np.zeros(end - start, dtype=np.float32)
        available_start = max(source_start, self.base)
        available_end = min(source_end, self.samples)
        if available_end > available_start:
            destination = profile.left_samples + available_start - start
            output[destination:destination + available_end - available_start] = self.buffer[
                available_start - self.base:available_end - self.base]
        return output


class TextDecoder:
    def __init__(self, table_path):
        data = Path(table_path).read_bytes()
        require(168 <= len(data) <= 16 * 1024 * 1024 and data[:8] == b'A8TOK001', 'tokenizer format/size')
        count, model_count, byte_count, *special = struct.unpack_from('<8I', data, 8)
        end = 168 + (count + 1) * 4
        require(5 <= count <= model_count <= 262144 and end + count + byte_count == len(data), 'tokenizer dimensions')
        offsets = struct.unpack_from(f'<{count + 1}I', data, 168)
        flags, pieces = data[end:end + count], data[end + count:]
        require(offsets[0] == 0 and offsets[-1] == byte_count and
                all(0 < b - a <= 4096 for a, b in zip(offsets, offsets[1:])), 'tokenizer offsets')
        require(set(flags) <= {0, 1} and len(set(special)) == 5 and
                all(i < count and flags[i] for i in special), 'tokenizer special IDs')
        self.count, self.model_count = count, model_count
        self.offsets, self.flags, self.pieces, self.special = offsets, flags, pieces, special
        self.decoder = codecs.getincrementaldecoder('utf-8')('replace')
        self.finished = False
        self.visible_tokens = self.text_bytes = 0
        self.table_sha256 = hashlib.sha256(data).hexdigest()

    def fresh(self):
        """Share the immutable tokenizer table, with a fresh UTF-8 decoder."""
        new = object.__new__(type(self))
        for name in ('count', 'model_count', 'offsets', 'flags', 'pieces', 'special', 'table_sha256'):
            setattr(new, name, getattr(self, name))
        new.decoder = codecs.getincrementaldecoder('utf-8')('replace')
        new.finished = False
        new.visible_tokens = new.text_bytes = 0
        return new

    def push(self, token):
        require(type(token) is int and 0 <= token < self.model_count, 'generated token outside vocabulary')
        if self.finished:
            return ''
        if token == self.special[3]:
            return self.finish()
        if token in self.special or token >= self.count or self.flags[token]:
            return ''
        self.visible_tokens += 1
        a, b = self.offsets[token:token + 2]
        text = self.decoder.decode(self.pieces[a:b], final=False)
        self.text_bytes += len(text.encode('utf-8'))
        return text

    def finish(self):
        if self.finished: return ''
        text = self.decoder.decode(b'', final=True)
        self.finished = True
        self.text_bytes += len(text.encode('utf-8'))
        return text