File size: 12,393 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
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
"""Actual sequential PCM streaming CLI; JSONL events, readable deltas flushed."""
from __future__ import annotations
import argparse
import hashlib
import importlib.metadata
import json
from pathlib import Path
import resource
import struct
import sys
import time
import numpy as np
import mlx.core as mx
from .frontend import Frontend
from .model import Audio8Model
from .stream import PCMReader, StreamProfile, TextDecoder
from .weights import sha256


def emit(kind, **fields):
    print(json.dumps({'kind': kind, **fields}, ensure_ascii=False, allow_nan=False), flush=True)


def memory():
    return {'mlx_active_bytes': mx.get_active_memory(), 'mlx_cache_bytes': mx.get_cache_memory(),
            'mlx_peak_bytes': mx.get_peak_memory(),
            'process_peak_rss_bytes': resource.getrusage(resource.RUSAGE_SELF).ru_maxrss * (1 if sys.platform == 'darwin' else 1024)}


def main(argv=None):
    parser = argparse.ArgumentParser(description='Experimental complete packed Audio8 MLX streaming recognizer')
    parser.add_argument('--bundle', type=Path, required=True)
    parser.add_argument('--config', type=Path, required=True)
    parser.add_argument('--frontend', type=Path, required=True)
    parser.add_argument('--tokenizer', type=Path, required=True)
    parser.add_argument('--silero-vad', type=Path, help='Causal silence resets with the pinned Silero ONNX model')
    parser.add_argument('--input', default='-', help='Raw finite mono 16 kHz F32LE; - means blocking stdin')
    parser.add_argument('--language', choices=['en', 'zh'], default='en')
    parser.add_argument('--startup', choices=['hf', 'realtime18'], default='realtime18')
    parser.add_argument('--delay-ms', type=int, choices=list(range(80, 1201, 80)) + [1600, 2400], default=240)
    parser.add_argument('--fuse-projections', action='store_true', help='Lossless compatible QKV/gate-up row fusion; experimental')
    parser.add_argument('--math-mode', choices=['reference', 'shared-rope', 'compiled'], default='reference',
        help='Optional shared rotary factors / pure compiled projection blocks; separate numerical gate required')
    parser.add_argument('--warmup-projections', action='store_true',
        help='Compile pure projection shapes before stream readiness; requires compiled math')
    parser.add_argument('--batch-windows', type=int, choices=[1, 2, 4], default=1, help='First batch1; later encoder windows, split before rolling')
    parser.add_argument('--dtype', choices=['float32', 'float16', 'bfloat16'], default='float32')
    parser.add_argument('--cache-dtype', choices=['float32', 'float16', 'bfloat16'], default='bfloat16')
    parser.add_argument('--pace-audio', action='store_true', help='Simulate file sample arrival; never a microphone claim')
    parser.add_argument('--max-emissions', type=int, default=0, help='Diagnostic bound; 0 reads through EOF')
    parser.add_argument('--dump-first-logits', type=Path, help='Exclusive .npy diagnostic, no overwrite')
    parser.add_argument('--memory-limit-gib', type=float, default=4.0, help='MLX allocator limit, not host RSS guarantee')
    args = parser.parse_args(argv)
    if args.max_emissions < 0 or not 1 <= args.memory_limit_gib <= 20:
        parser.error('invalid diagnostic/memory bound')
    if args.pace_audio and args.input == '-':
        parser.error('stdin is paced by its producer; --pace-audio is file only')
    if args.dump_first_logits and args.dump_first_logits.exists():
        parser.error('refusing to overwrite first logits')
    if args.warmup_projections and args.math_mode != 'compiled':
        parser.error('--warmup-projections requires --math-mode compiled')
    if args.silero_vad and (args.startup != 'realtime18' or args.delay_ms != 240 or args.batch_windows not in (1, 2, 4)
            or args.dtype != 'float16' or args.cache_dtype != 'bfloat16' or args.math_mode != 'compiled'
            or not args.fuse_projections or not args.warmup_projections or args.max_emissions or args.dump_first_logits):
        parser.error('--silero-vad requires realtime18/240ms/B1|2|4/float16/BF16/compiled with fused warmed projections and no diagnostic bounds')
    mx.set_default_device(mx.gpu)
    mx.set_memory_limit(int(args.memory_limit_gib * 1024 ** 3))
    mx.set_cache_limit(64 * 1024 * 1024)
    mx.reset_peak_memory()
    load_start = time.monotonic()
    frontend = Frontend.from_fixture(args.frontend)
    tokenizer = TextDecoder(args.tokenizer)
    profile = StreamProfile(args.startup, args.delay_ms // 80, 151668 if args.language == 'en' else 151667)
    model = Audio8Model.load(args.bundle, args.config, frontend, dtype=args.dtype, cache_dtype=args.cache_dtype,
        fuse_projections=args.fuse_projections, math_mode=args.math_mode)
    if tokenizer.model_count != model.text['vocab_size']:
        raise ValueError('tokenizer/model vocabulary mismatch')
    warmup_start = time.monotonic()
    warmup_calls = (model.warm_compiled_projections(profile.prefill, args.batch_windows)
                    if args.warmup_projections else 0)
    warmup_seconds = time.monotonic() - warmup_start if args.warmup_projections else 0.0
    if args.silero_vad:
        from .silero_stream import run
        return run(args, model, frontend, profile, load_start, warmup_seconds, warmup_calls)
    session = model.session(delay_tokens=profile.delay_tokens)
    mx.synchronize()
    source_hashes = {p.name: sha256(p) for p in sorted(Path(__file__).parent.glob('*.py'))}
    emit('mlx_stream_ready', backend='mlx_metal', mlx=importlib.metadata.version('mlx'),
        weights=model.weights.provenance, weight_resident_bytes=model.weights.nbytes,
        config_sha256=sha256(args.config), frontend=frontend.provenance, tokenizer_sha256=tokenizer.table_sha256,
        source_sha256=source_hashes, startup=args.startup, initial_prompt_tokens=profile.prefill,
        initial_generated_feedback=(args.startup == 'realtime18'), transcription_delay_ms=args.delay_ms,
        eof_policy=profile.eof_policy, dtype=args.dtype, cache_dtype=args.cache_dtype, math_mode=args.math_mode,
        projection_warmup=args.warmup_projections, projection_warmup_calls=warmup_calls,
        projection_warmup_seconds=warmup_seconds,
        activation_quantization=False, gear=4, samples_per_token=1280, token_period_ms=80,
        rolling_context=375, rolling_trim=38, rolling_stable_prefix=16,
        first_batch_windows=1, later_batch_windows=args.batch_windows,
        attention_history_rounding='per_original_window', fused_projection_groups=model.fused_projection_groups,
        input_mode='paced_file' if args.pace_audio else ('stdin' if args.input == '-' else 'unpaced_file'),
        retains_transcript=False, retains_generated_ids=False, full_model_quality_validated=False,
        load_seconds=time.monotonic() - load_start, memory=memory())
    stream = sys.stdin.buffer if args.input == '-' else open(args.input, 'rb')
    reader = PCMReader(stream)
    start = time.monotonic()
    ids_sha = hashlib.sha256()
    index, previous = 0, None
    first_text = last_text = None
    total_work = total_read = total_sleep = 0.0
    try:
        batch_index = 0
        while not args.max_emissions or index < args.max_emissions:
            inputs = profile.prompt() if index == 0 else [previous]
            requested = 1 if index == 0 else args.batch_windows
            if args.max_emissions: requested = min(requested, args.max_emissions - index)
            count = session.batch_capacity(requested, len(inputs))
            waveforms = []
            read_start = time.monotonic()
            for offset in range(count):
                waveform = reader.window(profile, index + offset)
                if waveform is None: break
                waveforms.append(waveform)
            total_read += time.monotonic() - read_start
            if not waveforms: break
            last_index = index + len(waveforms) - 1
            source_samples = min(profile.source_needed(last_index), reader.samples) if reader.eof else profile.source_needed(last_index)
            deadline = source_samples / 16000
            if args.pace_audio:
                duration = max(0., start + deadline - time.monotonic())
                if duration:
                    wait_start = time.monotonic(); time.sleep(duration); total_sleep += time.monotonic() - wait_start
            work_start = time.monotonic()
            queue = max(0., work_start - start - deadline) if args.pace_audio or args.input == '-' else None
            emission_start = work_start
            def token_ready(offset, output):
                nonlocal first_text, last_text, previous, emission_start
                if isinstance(output, tuple):
                    token, logits = output
                    with args.dump_first_logits.open('xb') as f:
                        np.save(f, np.array(logits), allow_pickle=False)
                else: token = output
                emitted_at = time.monotonic()
                interval = emitted_at - emission_start
                emission_start = emitted_at
                emitted = emitted_at - start
                ids_sha.update(struct.pack('<I', token))
                text = tokenizer.push(token)
                emit('mlx_stream_window', index=index + offset, token_id=token, audio_clock_seconds=emitted,
                    nominal_source_dependency_seconds=profile.source_needed(index + offset) / 16000,
                    available_batch_source_seconds=deadline, queue_seconds=queue,
                    work_interval_seconds=interval, batch_index=batch_index,
                    local_text_position=session.position, rolling_trims=session.trims,
                    encoder_cached_frames=session.encoder[0].length, decoder_cached_tokens=session.decoder[0].length,
                    kv_logical_bytes=session.cache_bytes, retained_input_samples=len(reader.buffer), memory=memory())
                if text:
                    first_text = emitted if first_text is None else first_text
                    last_text = emitted
                    emit('mlx_text_delta', index=index + offset, text=text, audio_clock_seconds=emitted)
                previous = token
            session.step_batch(waveforms, inputs,
                return_logits=index == 0 and args.dump_first_logits is not None, on_token=token_ready)
            work = time.monotonic() - work_start; total_work += work
            emit('mlx_stream_batch', batch_index=batch_index, first_window=index, windows=len(waveforms),
                work_seconds=work, source_ready_seconds=deadline, queue_seconds=queue,
                complete_source_batch=(index > 0 and len(waveforms) == args.batch_windows
                                       and profile.source_needed(last_index) <= reader.samples))
            index += len(waveforms); batch_index += 1
        tail = tokenizer.finish()
        if tail:
            emitted = time.monotonic() - start
            first_text = emitted if first_text is None else first_text; last_text = emitted
            emit('mlx_text_delta', index=index, text=tail, audio_clock_seconds=emitted, utf8_final_flush=True)
        complete = reader.eof and index == profile.emissions_at_eof(reader.samples)
        emit('mlx_stream_complete', status='complete' if complete else 'diagnostic_emission_limit',
            source_samples_observed=reader.samples, input_sha256_observed=reader.digest.hexdigest(),
            all_source_and_eof_windows_observed=complete, generated_tokens=index, generated_ids_sha256=ids_sha.hexdigest(),
            readable_text_bytes=tokenizer.text_bytes, first_text_audio_clock_seconds=first_text,
            last_text_audio_clock_seconds=last_text, configured_delay_ms=args.delay_ms,
            audio_seconds=reader.samples / 16000, elapsed_seconds=time.monotonic() - start,
            model_work_seconds=total_work, blocking_read_seconds=total_read, intentional_wait_seconds=total_sleep,
            max_retained_input_samples=reader.max_retained_samples, final_kv_logical_bytes=session.cache_bytes,
            batch_count=batch_index, first_batch_windows=1, later_batch_windows=args.batch_windows,
            rolling_trims=session.trims, memory=memory(),
            measured_word_aligned_latency=False, production_accuracy_approved=False)
    finally:
        if args.input != '-': stream.close()


if __name__ == '__main__':
    main()