Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
52 changes: 51 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -217,13 +217,14 @@ audio = model.generate(text="He plays the [B EY1 S] guitar while catching a [B A

## Command-Line Tools

Three CLI entry points are provided. The CLI tools support all features available in the Python API (voice cloning, voice design, auto voice, generation parameters, etc.) — all controlled via command-line arguments.
Four CLI entry points are provided. The CLI tools support all features available in the Python API (voice cloning, voice design, auto voice, generation parameters, etc.) — all controlled via command-line arguments.

| Command | Description | Source |
|---|---|---|
| `omnivoice-demo` | Interactive Gradio web demo | [omnivoice/cli/demo.py](omnivoice/cli/demo.py) |
| `omnivoice-infer` | Single-item inference | [omnivoice/cli/infer.py](omnivoice/cli/infer.py) |
| `omnivoice-infer-batch` | Batch inference across multiple GPUs | [omnivoice/cli/infer_batch.py](omnivoice/cli/infer_batch.py) |
| `omnivoice-infer-mlx` | Single-item MLX inference on Apple Silicon | [omnivoice/cli/infer_mlx.py](omnivoice/cli/infer_mlx.py) |

### Demo

Expand Down Expand Up @@ -258,6 +259,55 @@ omnivoice-infer \
--output hello.wav
```

### MLX Inference on Apple Silicon

An experimental MLX backend is available for Apple Silicon. It runs the
OmniVoice diffusion language model with MLX and keeps the Higgs audio tokenizer
on Transformers/PyTorch.

```bash
pip install -e ".[mlx]"

omnivoice-infer-mlx \
--model k2-fsa/OmniVoice \
--text "This is a test for MLX inference." \
--instruct "female, british accent" \
--output hello_mlx.wav
```

Python API:

```python
from omnivoice.mlx import OmniVoiceMLX
import soundfile as sf

model = OmniVoiceMLX.from_pretrained("k2-fsa/OmniVoice", dtype="float16")
audio = model.generate(
text="Hello from the MLX backend.",
instruct="female, british accent",
)
sf.write("out.wav", audio[0], model.sampling_rate)
```

### MLX Conversion and Staging

The repository includes helper scripts for reproducible local MLX exports and
future Hugging Face uploads:

```bash
# Build five local staging directories:
# OmniVoice, OmniVoice-fp32, OmniVoice-bf16,
# OmniVoice-8bit, OmniVoice-4bit
python scripts/stage_mlx_repos.py \
--source /path/to/OmniVoice-official \
--output-root /path/to/huggingface

# Validate one staged directory.
python scripts/validate_mlx.py \
--model /path/to/huggingface/OmniVoice \
--output smoke.wav
```

### Batch Inference

`omnivoice-infer-batch` can distribute batch inference across multiple GPUs, designed for large-scale TTS tasks.
Expand Down
86 changes: 86 additions & 0 deletions omnivoice/cli/infer_mlx.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,86 @@
"""Single-item MLX inference CLI for OmniVoice."""

import argparse
import logging

import soundfile as sf

from omnivoice.mlx import OmniVoiceMLX
from omnivoice.utils.common import str2bool


def get_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description="OmniVoice single-item inference with the MLX backend",
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
parser.add_argument("--model", type=str, default="k2-fsa/OmniVoice")
parser.add_argument("--text", type=str, required=True)
parser.add_argument("--output", type=str, required=True)
parser.add_argument("--ref_audio", type=str, default=None)
parser.add_argument("--ref_text", type=str, default=None)
parser.add_argument("--instruct", type=str, default=None)
parser.add_argument("--language", type=str, default=None)
parser.add_argument("--num_step", type=int, default=32)
parser.add_argument("--guidance_scale", type=float, default=2.0)
parser.add_argument("--speed", type=float, default=1.0)
parser.add_argument("--duration", type=float, default=None)
parser.add_argument("--t_shift", type=float, default=0.1)
parser.add_argument("--denoise", type=str2bool, default=True)
parser.add_argument("--postprocess_output", type=str2bool, default=True)
parser.add_argument("--layer_penalty_factor", type=float, default=5.0)
parser.add_argument("--position_temperature", type=float, default=5.0)
parser.add_argument("--class_temperature", type=float, default=0.0)
parser.add_argument(
"--dtype",
type=str,
default="float16",
choices=["float16", "bfloat16", "float32"],
help="MLX weight and compute dtype.",
)
parser.add_argument(
"--audio_tokenizer_device",
type=str,
default="cpu",
help="Device map for the Transformers Higgs audio tokenizer.",
)
return parser


def main():
formatter = "%(asctime)s %(levelname)s [%(filename)s:%(lineno)d] %(message)s"
logging.basicConfig(format=formatter, level=logging.INFO, force=True)
args = get_parser().parse_args()

logging.info("Loading MLX model from %s ...", args.model)
model = OmniVoiceMLX.from_pretrained(
args.model,
dtype=args.dtype,
audio_tokenizer_device=args.audio_tokenizer_device,
)

logging.info("Generating audio for: %s...", args.text[:80])
audios = model.generate(
text=args.text,
language=args.language,
ref_audio=args.ref_audio,
ref_text=args.ref_text,
instruct=args.instruct,
duration=args.duration,
num_step=args.num_step,
guidance_scale=args.guidance_scale,
speed=args.speed,
t_shift=args.t_shift,
denoise=args.denoise,
postprocess_output=args.postprocess_output,
layer_penalty_factor=args.layer_penalty_factor,
position_temperature=args.position_temperature,
class_temperature=args.class_temperature,
)

sf.write(args.output, audios[0], model.sampling_rate)
logging.info("Saved to %s", args.output)


if __name__ == "__main__":
main()
5 changes: 5 additions & 0 deletions omnivoice/mlx/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
"""MLX inference backend for OmniVoice."""

from omnivoice.mlx.omnivoice import OmniVoiceMLX

__all__ = ["OmniVoiceMLX"]
Loading