CoolFace
Modelpublic

fredchu/MOSS-Audio-8B-Instruct-MLX

sourceHugging Faceapache-2.0updated 4mo agoView on Hugging Face
3likes
inference.py108 linesDownload Raw Back to root
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