Reza2kn's picture
Publish complete experimental Audio8 Q4 model and runtimes
21fd722 verified
Raw History Blame Contribute Delete
7.37 kB
"""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