Automatic Speech Recognition
MLX
ONNX
GGUF
Rust
English
Chinese
audio8
streaming-asr
quantized
experimental
Instructions to use Reza2kn/Audio8-ASR-Infinite-Compressed with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use Reza2kn/Audio8-ASR-Infinite-Compressed with MLX:
# Download the model from the Hub pip install huggingface_hub[hf_xet] huggingface-cli download --local-dir Audio8-ASR-Infinite-Compressed Reza2kn/Audio8-ASR-Infinite-Compressed
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Atomic Chat
Download mlx_runtime/model.py from Reza2kn/Audio8-ASR-Infinite-Compressed: direct link, hf CLI and curl.
- Browser
- Download file 17.9 kB
-
https://huggingface.co/Reza2kn/Audio8-ASR-Infinite-Compressed/resolve/main/mlx_runtime/model.py
- Command line
-
hf download hf://Reza2kn/Audio8-ASR-Infinite-Compressed/mlx_runtime/model.py
-
curl -L -o model.py https://huggingface.co/Reza2kn/Audio8-ASR-Infinite-Compressed/resolve/main/mlx_runtime/model.py
17.9 kB
| """Complete Audio8 encoder/projector/conditioned Qwen2 decoder and tied head.""" | |
| from __future__ import annotations | |
| import json | |
| from pathlib import Path | |
| import mlx.core as mx | |
| from .weights import Weights, require, sha256 | |
| from .math import KVCache, attention, gelu, rms_norm, fast_rms_norm, rope, rotary_factors, apply_rotary | |
| DTYPES = {'float32': mx.float32, 'float16': mx.float16, 'bfloat16': mx.bfloat16} | |
| class Audio8Model: | |
| def __init__(self, weights, config, frontend, *, dtype='float32', cache_dtype='bfloat16', fuse_projections=False, math_mode='reference'): | |
| require(config['text_config']['model_type'] == 'qwen2', 'only pinned Qwen2 text architecture') | |
| require(config['max_frame_len'] == 8 and config['frame_lens'] == [4, 6, 8], 'projector/frame configuration') | |
| self.weights, self.config, self.frontend = weights, config, frontend | |
| self.dtype, self.cache_dtype = DTYPES[dtype], DTYPES[cache_dtype] | |
| self.dtype_name, self.cache_dtype_name = dtype, cache_dtype | |
| require(math_mode in ('reference', 'shared-rope', 'compiled'), 'math mode') | |
| self.math_mode = math_mode | |
| self.compiled_blocks = {} | |
| self.audio = config['audio_config']; self.text = config['text_config'] | |
| self.audio_theta = self.audio['rope_parameters']['rope_theta'] | |
| self.text_theta = self.text['rope_parameters']['rope_theta'] | |
| for i in range(self.audio['num_hidden_layers']): | |
| self._validate_layer(f'audio_tower.layers.{i}', True) | |
| for i in range(self.text['num_hidden_layers']): | |
| self._validate_layer(f'language_model.model.layers.{i}', False) | |
| self.fused_projection_groups = 0 | |
| if fuse_projections: | |
| for stem, layers in [('audio_tower.layers', self.audio['num_hidden_layers']), | |
| ('language_model.model.layers', self.text['num_hidden_layers'])]: | |
| for index in range(layers): | |
| p = f'{stem}.{index}' | |
| for suffixes in [('self_attn.q_proj', 'self_attn.k_proj', 'self_attn.v_proj'), | |
| ('mlp.gate_proj', 'mlp.up_proj')]: | |
| self.fused_projection_groups += weights.fuse([p + '.' + suffix for suffix in suffixes]) | |
| head = weights.tensors['language_model.model.embed_tokens.weight'] | |
| require((head.rows, head.cols) == (self.text['vocab_size'], self.text['hidden_size']), 'embedding dimensions') | |
| def _validate_layer(self, prefix, encoder): | |
| config = self.audio if encoder else self.text | |
| h, n, intermediate = config['hidden_size'], config['num_attention_heads'], config['intermediate_size'] | |
| dim = config.get('head_dim', h // n) | |
| kv = config.get('num_key_value_heads', n) | |
| for sub, shape in {'self_attn.q_proj': (n * dim, h), 'self_attn.k_proj': (kv * dim, h), | |
| 'self_attn.v_proj': (kv * dim, h), 'self_attn.o_proj': (h, n * dim), | |
| 'mlp.gate_proj': (intermediate, h), 'mlp.up_proj': (intermediate, h), | |
| 'mlp.down_proj': (h, intermediate)}.items(): | |
| weight = self.weights.tensors[prefix + '.' + sub + '.weight'] | |
| require((weight.rows, weight.cols) == shape, f'layer weight shape: {prefix}.{sub}') | |
| def load(cls, bundle, config_path, frontend, **kwargs): | |
| config = json.loads(Path(config_path).read_text()) | |
| weights = Weights.load(bundle) | |
| weights.provenance['config_sha256'] = sha256(config_path) | |
| require(weights.provenance['expected_config_sha256'] == weights.provenance['config_sha256'], 'config differs from bundle source asset') | |
| return cls(weights, config, frontend, **kwargs) | |
| def session(self, *, gear=4, delay_tokens=3, context=375, trim=38, stable=16): | |
| return Session(self, gear, delay_tokens, context, trim, stable) | |
| def warm_compiled_projections(self, prefill, batch_windows): | |
| """Compile bounded projection shapes before audio readiness. | |
| Only pure projection functions receive synthetic inputs. No session, | |
| tokenizer, audio reader, position, or KV cache is created or advanced. | |
| Explicit weights are shared with subsequent inference unchanged. | |
| """ | |
| require(self.math_mode == 'compiled', 'projection warmup requires compiled math') | |
| require(type(prefill) is int and 1 <= prefill <= 64, 'warmup prefill') | |
| require(type(batch_windows) is int and batch_windows in (1, 2, 4), 'warmup batch') | |
| calls = 0 | |
| for encoder, config, stem, lengths in ( | |
| (True, self.audio, 'audio_tower.layers', sorted({4 * prefill, *range(4, 4 * batch_windows + 1, 4)})), | |
| (False, self.text, 'language_model.model.layers', sorted({1, prefill})), | |
| ): | |
| modulation = (mx.array(1, self.dtype) if encoder | |
| else mx.ones((1, config['hidden_size']), self.dtype)) | |
| for index in range(config['num_hidden_layers']): | |
| (qkv, qp), (ffn, fp) = self._projection_blocks(f'{stem}.{index}', encoder) | |
| for length in lengths: | |
| x = mx.zeros((1, length, config['hidden_size']), self.dtype) | |
| mx.eval(*qkv(x, qp), ffn(x, modulation, fp)) | |
| calls += 2 | |
| mx.synchronize() | |
| return calls | |
| def _projection_blocks(self, prefix, encoder): | |
| """Pure stateless functions: no cache/position mutation or array copies. | |
| Hidden dimensions, weight identity, epsilon and norm dtype are fixed. | |
| Two function objects per layer cache only bounded token-length shapes: | |
| at most the initial encoder length plus4/8/12/16, and initial/one-token | |
| decoder lengths. No position/history length enters these functions. | |
| """ | |
| if prefix not in self.compiled_blocks: | |
| cfg = self.audio if encoder else self.text | |
| norm1 = 'self_attn_layer_norm' if encoder else 'input_layernorm' | |
| norm2 = 'final_layer_norm' if encoder else 'post_attention_layernorm' | |
| weights, eps = self.weights, cfg['rms_norm_eps'] | |
| norm1_weight = weights.tensor(prefix + '.' + norm1 + '.weight') | |
| norm2_weight = weights.tensor(prefix + '.' + norm2 + '.weight') | |
| qkv_names = [prefix + '.self_attn.' + name for name in ('q_proj', 'k_proj', 'v_proj')] | |
| gate_up_names = [prefix + '.mlp.gate_proj', prefix + '.mlp.up_proj'] | |
| qkv_op, qkv_arrays = weights.linear_plan(qkv_names) | |
| gate_up_op, gate_up_arrays = weights.linear_plan(gate_up_names) | |
| down_op, down_arrays = weights.linear_plan([prefix + '.mlp.down_proj']) | |
| def qkv(x, parameters): | |
| z = fast_rms_norm(x, parameters['norm'], eps) | |
| return qkv_op(z, parameters['projection']) | |
| def ffn(x, modulation, parameters): | |
| z = fast_rms_norm(x, parameters['norm'], eps) * modulation | |
| gate, up = gate_up_op(z, parameters['gate_up']) | |
| return x + down_op((gate * mx.sigmoid(gate)) * up, parameters['down'])[0] | |
| # Explicit function arguments, rather than closure constants or | |
| # a mutable captured list, preserve exact shared weight ownership. | |
| qkv_inputs = {'norm': norm1_weight, 'projection': qkv_arrays} | |
| ffn_inputs = {'norm': norm2_weight, 'gate_up': gate_up_arrays, 'down': down_arrays} | |
| # MLX0.32 cannot infer fused-output slice shapes in shapeless mode | |
| # when packed arrays are explicit arguments. Bounded fixed shapes | |
| # are valid here; unlike whole layers, no growing KV shape enters. | |
| self.compiled_blocks[prefix] = ((mx.compile(qkv), qkv_inputs), | |
| (mx.compile(ffn), ffn_inputs)) | |
| return self.compiled_blocks[prefix] | |
| def layer(self, x, cache, prefix, position, encoder, modulation=None, attention_partitions=None, rotary=None): | |
| cfg = self.audio if encoder else self.text | |
| heads = cfg['num_attention_heads']; dim = cfg.get('head_dim', cfg['hidden_size'] // heads) | |
| kv_heads = cfg.get('num_key_value_heads', heads) | |
| norm1 = 'self_attn_layer_norm' if encoder else 'input_layernorm' | |
| norm2 = 'final_layer_norm' if encoder else 'post_attention_layernorm' | |
| w = self.weights | |
| compiled = getattr(self, 'math_mode', 'reference') == 'compiled' | |
| if compiled: | |
| block, parameters = self._projection_blocks(prefix, encoder)[0] | |
| q, k, v = block(x, parameters) | |
| else: | |
| z = rms_norm(x, w.tensor(prefix + '.' + norm1 + '.weight'), cfg['rms_norm_eps']) | |
| q, k, v = w.linear_many([prefix + '.self_attn.' + name for name in ('q_proj', 'k_proj', 'v_proj')], z) | |
| q = q.reshape(1, -1, heads, dim).transpose(0, 2, 1, 3) | |
| k = k.reshape(1, -1, kv_heads, dim).transpose(0, 2, 1, 3) | |
| v = v.reshape(1, -1, kv_heads, dim).transpose(0, 2, 1, 3) | |
| theta = self.audio_theta if encoder else self.text_theta | |
| q, k = ((rope(q, position, theta), rope(k, position, theta)) if rotary is None | |
| else (apply_rotary(q, rotary), apply_rotary(k, rotary))) | |
| if attention_partitions is None: | |
| result = attention(q, k, v, cache, window=cfg['sliding_window'] if encoder else None) | |
| else: | |
| require(encoder and sum(attention_partitions) == q.shape[2] | |
| and all(n > 0 for n in attention_partitions), 'attention partitions') | |
| contexts, offset = [], 0 | |
| for count in attention_partitions: | |
| stop = offset + count | |
| contexts.append(attention(q[:, :, offset:stop], k[:, :, offset:stop], | |
| v[:, :, offset:stop], cache, window=cfg['sliding_window'])) | |
| offset = stop | |
| result = mx.concatenate(contexts, axis=2) | |
| x = x + w.linear(prefix + '.self_attn.o_proj', result.transpose(0, 2, 1, 3).reshape(1, -1, heads * dim)) | |
| if compiled: | |
| modulation = mx.array(1, x.dtype) if modulation is None else modulation | |
| block, parameters = self._projection_blocks(prefix, encoder)[1] | |
| return block(x, modulation, parameters) | |
| z = rms_norm(x, w.tensor(prefix + '.' + norm2 + '.weight'), cfg['rms_norm_eps']) | |
| if modulation is not None: | |
| z = z * modulation | |
| gate, up = w.linear_many([prefix + '.mlp.gate_proj', prefix + '.mlp.up_proj'], z) | |
| return x + w.linear(prefix + '.mlp.down_proj', (gate * mx.sigmoid(gate)) * up) | |
| class Session: | |
| def __init__(self, model, gear, delay_tokens, context, trim, stable): | |
| require(gear in model.config['frame_lens'] and type(delay_tokens) is int and 1 <= delay_tokens <= 30, 'gear/delay') | |
| require(context > stable + trim > stable > 1 and trim > 0, 'rolling policy') | |
| self.model, self.gear, self.delay_tokens = model, gear, delay_tokens | |
| self.context, self.trim, self.stable = context, trim, stable | |
| self.position = self.emissions = self.trims = 0 | |
| self.encoder = [KVCache(model.audio['sliding_window'] - 1, model.cache_dtype, sliding=True) | |
| for _ in range(model.audio['num_hidden_layers'])] | |
| self.decoder = [KVCache(context, model.cache_dtype) for _ in range(model.text['num_hidden_layers'])] | |
| hidden = model.text['hidden_size'] | |
| inv = mx.exp(-mx.log(mx.array(10000., mx.float32)) * mx.arange(hidden // 2, dtype=mx.float32) / (hidden // 2)).astype(model.dtype) | |
| phase = mx.array(delay_tokens, model.dtype) * inv | |
| condition = mx.concatenate([mx.cos(phase), mx.sin(phase)]) | |
| condition = condition + model.weights.tensor('frame_len_embedding.weight')[model.config['frame_lens'].index(gear)].astype(model.dtype) | |
| self.modulations = [] | |
| for i in range(model.text['num_hidden_layers']): | |
| prefix = f'language_model.model.layers.{i}.ada_rms_norm' | |
| value = model.weights.linear(prefix + '.linear2', gelu(model.weights.linear(prefix + '.linear1', condition[None]))) | |
| self.modulations.append((1 + value).astype(model.dtype)) | |
| mx.eval(*self.modulations) | |
| def maybe_trim(self, incoming): | |
| require(0 < incoming <= self.context - self.stable, 'incoming token count') | |
| while self.position + incoming > self.context: | |
| require(all(c.length == self.position for c in self.decoder), 'decoder cache clock mismatch') | |
| for cache in self.decoder: | |
| cache.trim_decoder(self.trim, self.stable, self.model.text_theta) | |
| for cache in self.encoder: | |
| cache.rebase_encoder(self.trim * self.gear, self.model.audio_theta) | |
| # Bound rotation scratch to one layer; preserve the exact operations. | |
| mx.eval(*cache.arrays()) | |
| self.position -= self.trim; self.trims += 1 | |
| mx.eval(*self.cache_arrays()) | |
| def cache_arrays(self): | |
| return [a for c in self.encoder + self.decoder for a in c.arrays()] | |
| def cache_bytes(self): | |
| return sum(c.nbytes for c in self.encoder + self.decoder) | |
| def batch_capacity(self, requested, initial_tokens=1): | |
| """Trim before the first query if required, then never cross another trim.""" | |
| require(type(requested) is int and 1 <= requested <= 4, 'batch windows must be1..4') | |
| require(1 <= initial_tokens <= 64, 'initial token count') | |
| position = self.position | |
| while position + initial_tokens > self.context: | |
| position -= self.trim | |
| return min(requested, 1 + self.context - position - initial_tokens) | |
| def encode_windows(self, waveforms, counts): | |
| m, w = self.model, self.model.weights | |
| parts = [m.frontend.convolve(m.frontend.mel(wave), w, count * self.gear, m.dtype) | |
| for wave, count in zip(waveforms, counts)] | |
| x = mx.concatenate(parts, axis=1) if len(parts) > 1 else parts[0] | |
| partitions = [count * self.gear for count in counts] | |
| rotary = None | |
| if getattr(m, 'math_mode', 'reference') != 'reference': | |
| rotary = rotary_factors(self.position * self.gear, x.shape[1], | |
| m.audio.get('head_dim', m.audio['hidden_size'] // m.audio['num_attention_heads']), m.dtype, m.audio_theta) | |
| for i, cache in enumerate(self.encoder): | |
| x = m.layer(x, cache, f'audio_tower.layers.{i}', self.position * self.gear, | |
| True, attention_partitions=partitions, rotary=rotary) | |
| x = rms_norm(x, w.tensor('audio_tower.norm.weight'), m.audio['rms_norm_eps']) | |
| grouped = x.reshape(1, sum(counts), self.gear, m.audio['hidden_size']) | |
| grouped = mx.pad(grouped, [(0, 0), (0, 0), (0, 8 - self.gear), (0, 0)]).reshape(1, sum(counts), -1) | |
| audio = w.linear('multi_modal_projector.linear_2', gelu(w.linear('multi_modal_projector.linear_1', grouped))) | |
| # Finish the encoder batch before autoregressive feedback. Cache history | |
| # was rounded at each original window, not only at the batch endpoint. | |
| mx.eval(audio, *[a for c in self.encoder for a in c.arrays()]) | |
| return audio | |
| def decode_audio(self, audio, token_ids, *, return_logits=False): | |
| m, w = self.model, self.model.weights | |
| x = w.tensors['language_model.model.embed_tokens.weight'].embedding(token_ids, m.dtype)[None] + audio | |
| rotary = None | |
| if getattr(m, 'math_mode', 'reference') != 'reference': | |
| rotary = rotary_factors(self.position, len(token_ids), | |
| m.text.get('head_dim', m.text['hidden_size'] // m.text['num_attention_heads']), m.dtype, m.text_theta) | |
| for i, cache in enumerate(self.decoder): | |
| x = m.layer(x, cache, f'language_model.model.layers.{i}', self.position, False, self.modulations[i], rotary=rotary) | |
| x = rms_norm(x[:, -1:], w.tensor('language_model.model.norm.weight'), m.text['rms_norm_eps']) | |
| logits = w.tensors['language_model.lm_head.weight'](x)[0, 0].astype(mx.float32) | |
| logits = mx.where(mx.arange(logits.size) == m.config['eos_token_id'], -float('inf'), logits) | |
| choice = mx.argmax(logits) | |
| maximum = mx.max(logits) | |
| mx.eval(choice, maximum, *[a for c in self.decoder for a in c.arrays()]) | |
| require(bool(mx.isfinite(maximum).item()), 'nonfinite output logits') | |
| token = int(choice.item()) | |
| self.position += len(token_ids); self.emissions += 1 | |
| if return_logits: mx.eval(logits) | |
| return (token, logits) if return_logits else token | |
| def step_batch(self, waveforms, token_ids, *, return_logits=False, on_token=None): | |
| """Batch encoder projections; decode feedback sequentially; no trim inside. | |
| At most four source windows. Callback fires after each actual decoder | |
| result, so no artificial final-batch timestamp is assigned to early text. | |
| """ | |
| require(1 <= len(waveforms) <= 4 and 1 <= len(token_ids) <= 64, 'batch shape') | |
| self.maybe_trim(len(token_ids)) | |
| counts = [len(token_ids)] + [1] * (len(waveforms) - 1) | |
| require(self.position + sum(counts) <= self.context, 'batch crosses a rolling boundary') | |
| audio = self.encode_windows(waveforms, counts) | |
| outputs, offset, inputs = [], 0, token_ids | |
| for number, count in enumerate(counts): | |
| output = self.decode_audio(audio[:, offset:offset + count], inputs, | |
| return_logits=return_logits) | |
| outputs.append(output) | |
| token = output[0] if isinstance(output, tuple) else output | |
| if on_token is not None: on_token(number, output) | |
| inputs = [token]; offset += count | |
| return outputs | |
| def step(self, waveform, token_ids, *, return_logits=False): | |
| return self.step_batch([waveform], token_ids, return_logits=return_logits)[0] | |