epsilon3's picture
Make release M4-only and lead with speed and memory
728caeb verified
Raw History Blame Contribute Delete
1.33 kB
"""
Unified Engine Router for Parallel Constrained Decoding.
Automatically selects MLX backend on Apple Silicon macOS,
or the PyTorch / MPS backend on Apple Silicon.
"""
import os
import platform
USE_MLX = False
if platform.system() == "Darwin" and os.environ.get("BACKEND", "").lower() != "torch":
try:
import mlx.core as mx
import mlx_lm
USE_MLX = True
except Exception:
USE_MLX = False
if USE_MLX:
from core.engine_mlx import (
get_engine,
run_parallel_generation,
run_naive_generation,
stream_naive_generation,
run_rlcd_generation,
)
else:
from core.engine_torch import (
get_torch_engine as get_engine,
run_parallel_generation_torch as run_parallel_generation,
run_naive_generation_torch as run_naive_generation,
stream_naive_generation_torch as stream_naive_generation,
)
if os.environ.get("RLCD_ATTENTION", "tree").lower() != "batch":
from core.engine_tree import run_parallel_generation_tree as run_parallel_generation
run_rlcd_generation = run_parallel_generation
__all__ = [
"get_engine",
"run_parallel_generation",
"run_naive_generation",
"stream_naive_generation",
"run_rlcd_generation",
"USE_MLX",
]