fredchu/MOSS-Audio-8B-Instruct-MLX
3
1"""Standalone MOSS-Audio-{4B,8B}-Thinking MLX inference.2 3Usage:4 python inference.py --audio path/to/clip.wav [--max-tokens 2048]5 6Both 4B INT4 and 8B hybrid bundles work with this script. Audio-path7dtype is inferred from the saved adapter weights (`scales` key => INT4).8"""9from __future__ import annotations10import argparse, sys, time11from pathlib import Path12 13HERE = Path(__file__).resolve().parent14sys.path.insert(0, str(HERE / "scripts"))15 16import librosa17import mlx.core as mx18import numpy as np19from mlx_lm import load as mlx_load20from mlx_lm.generate import generate_step21from mlx_lm.sample_utils import make_sampler, make_logits_processors22 23from moss_audio_mlx_bridge_v3 import (24 load_mlx_audio_path,25 build_mel_spectrogram,26 run_mlx_audio_pipeline,27 install_deepstack_hooks,28)29 30 31def main():32 p = argparse.ArgumentParser()33 p.add_argument("--audio", required=True, help="Path to input .wav (16 kHz, mono)")34 p.add_argument("--max-tokens", type=int, default=2048)35 p.add_argument("--repetition-penalty", type=float, default=1.02,36 help="1.02 kills decode-loops without over-penalizing descriptions")37 args = p.parse_args()38 39 ad_w = mx.load(str(HERE / "mlx_audio/audio_adapter.safetensors"))40 if "down_proj.scales" in ad_w:41 llm_hidden = ad_w["down_proj.scales"].shape[0]42 int4_audio = True43 else:44 llm_hidden = ad_w["down_proj.weight"].shape[0]45 int4_audio = False46 size_tag = "4B" if llm_hidden == 2560 else "8B"47 print(f"[detect] {size_tag} bundle, audio int4={int4_audio}")48 49 print(f"[load] LLM from {HERE / 'mlx_llm'}")50 t0 = time.perf_counter()51 mlx_model, mlx_tokenizer = mlx_load(str(HERE / "mlx_llm"))52 print(f"[load] LLM: {time.perf_counter()-t0:.1f}s")53 54 t0 = time.perf_counter()55 encoder, adapter, mergers = load_mlx_audio_path(HERE / "mlx_audio", int4=int4_audio)56 print(f"[load] audio path: {time.perf_counter()-t0:.1f}s")57 58 y, _ = librosa.load(args.audio, sr=16000, mono=True)59 y = y.astype(np.float32)60 print(f"[audio] {args.audio} ({len(y)/16000:.1f}s)")61 62 # Pure-MLX mel + input_ids (no torch).63 mel, lens, input_ids_mx, audio_token_id = build_mel_spectrogram(y, mlx_tokenizer)64 primary, ds_embeds = run_mlx_audio_pipeline(encoder, adapter, mergers, mel, lens)65 primary = primary.astype(mx.bfloat16)66 ds_embeds = [d.astype(mx.bfloat16) for d in ds_embeds]67 mx.eval(primary, *ds_embeds)68 69 del encoder, adapter, mergers, mel, lens70 import gc; gc.collect(); mx.clear_cache()71 72 audio_mask = input_ids_mx == audio_token_id73 audio_positions = np.where(np.array(audio_mask[0]))[0]74 text_embeds = mlx_model.model.embed_tokens(input_ids_mx)75 text_np = np.array(text_embeds.astype(mx.float32))76 primary_np = np.array(primary.astype(mx.float32))77 text_np[0, audio_positions, :] = primary_np[0, :, :]78 merged = mx.array(text_np).astype(mx.bfloat16)79 80 ds_flat = [d[0] for d in ds_embeds]81 install_deepstack_hooks(mlx_model, ds_flat, audio_positions)82 83 sampler = make_sampler(temp=1.0, top_p=1.0, top_k=50)84 logits_processors = make_logits_processors(85 repetition_penalty=args.repetition_penalty, repetition_context_size=2086 ) if args.repetition_penalty else None87 88 gen_kwargs = dict(89 prompt=input_ids_mx[0], model=mlx_model,90 input_embeddings=merged[0], max_tokens=args.max_tokens, sampler=sampler,91 )92 if logits_processors:93 gen_kwargs["logits_processors"] = logits_processors94 95 t0 = time.perf_counter()96 generated = []97 for tok, _ in generate_step(**gen_kwargs):98 generated.append(int(tok))99 if tok == mlx_tokenizer.eos_token_id:100 break101 elapsed = time.perf_counter() - t0102 print(f"[gen] {len(generated)} tokens in {elapsed:.2f}s ({len(generated)/elapsed:.1f} t/s)")103 print(f"\n=== OUTPUT ===\n{mlx_tokenizer.decode(generated)}\n=== END ===")104 105 106if __name__ == "__main__":107 main()108 