"""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='= 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