Skip to content

Commit 16e3654

Browse files
feat(server): support custom model path and hub
Adds funasr-server --model-path/--hub support with maintainer follow-up for default Fun-ASR-Nano hub routing and custom fallback loading. Validated with py_compile, focused CLI/server tests, and git diff check.
1 parent dc9758d commit 16e3654

3 files changed

Lines changed: 112 additions & 14 deletions

File tree

funasr/bin/_server_app.py

Lines changed: 30 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -101,7 +101,7 @@ def prepare_audio_for_inference(audio_data, sr, target_sr=16000):
101101

102102
return audio_data.astype(np.float32), sr
103103

104-
def create_app(device: str = "cuda", preload_model: str = "auto") -> FastAPI:
104+
def create_app(device: str = "cuda", preload_model: str = "auto", model_path: str = None, hub: str = "ms") -> FastAPI:
105105
if preload_model == "auto":
106106
preload_model = "fun-asr-nano" if device.startswith("cuda") else "sensevoice"
107107

@@ -110,6 +110,8 @@ def create_app(device: str = "cuda", preload_model: str = "auto") -> FastAPI:
110110
app.state.engine = None
111111
app.state.vad_model = None
112112
app.state.fallback_models = {}
113+
app.state.model_path = model_path
114+
app.state.hub = hub
113115

114116
# Non-LLM model configs (use AutoModel, no vLLM)
115117
FALLBACK_CONFIGS = {
@@ -135,9 +137,12 @@ def _load_vllm_engine():
135137

136138
logger.info("Loading Fun-ASR-Nano vLLM engine...")
137139
t0 = time.time()
140+
# Use custom model_path if provided, otherwise default
141+
vllm_model = app.state.model_path if app.state.model_path else "FunAudioLLM/Fun-ASR-Nano-2512"
142+
vllm_hub = app.state.hub if app.state.model_path else "hf"
138143
app.state.engine = FunASRNanoVLLM.from_pretrained(
139-
model="FunAudioLLM/Fun-ASR-Nano-2512",
140-
hub="hf",
144+
model=vllm_model,
145+
hub=vllm_hub,
141146
device=device,
142147
dtype="bf16",
143148
max_model_len=4096,
@@ -154,28 +159,34 @@ def _load_vllm_engine():
154159
app.state.use_vllm = False
155160
from funasr import AutoModel
156161
cfg = {
157-
"model": "FunAudioLLM/Fun-ASR-Nano-2512",
158-
"hub": "hf",
162+
"model": app.state.model_path if app.state.model_path else "FunAudioLLM/Fun-ASR-Nano-2512",
163+
"hub": app.state.hub if app.state.model_path else "hf",
159164
"trust_remote_code": True,
160165
"vad_model": "fsmn-vad",
161166
"vad_kwargs": {"max_single_segment_time": 30000},
162167
"device": device,
163168
"disable_update": True,
164169
}
165170
app.state.fallback_models["fun-asr-nano"] = AutoModel(**cfg)
166-
logger.info("Fallback AutoModel loaded for fun-asr-nano.")
171+
logger.info(f"Fallback AutoModel loaded for fun-asr-nano with model={cfg['model']}, hub={cfg['hub']}.")
167172

168173
def _load_fallback(name: str):
169174
"""Load non-LLM model via AutoModel."""
170175
if name in app.state.fallback_models:
171176
return app.state.fallback_models[name]
172-
if name not in FALLBACK_CONFIGS:
177+
if name not in FALLBACK_CONFIGS and not app.state.model_path:
173178
return None
174179
from funasr import AutoModel
175-
cfg = FALLBACK_CONFIGS[name].copy()
180+
cfg = FALLBACK_CONFIGS.get(name, {}).copy()
181+
# Override with custom model_path and hub if provided
182+
if app.state.model_path:
183+
cfg["model"] = app.state.model_path
184+
cfg["hub"] = app.state.hub
185+
elif app.state.hub:
186+
cfg["hub"] = app.state.hub
176187
cfg["device"] = device
177188
cfg["disable_update"] = True
178-
logger.info(f"Loading fallback model '{name}'...")
189+
logger.info(f"Loading fallback model '{name}' with model={cfg['model']}, hub={cfg['hub']}...")
179190
model = AutoModel(**cfg)
180191
app.state.fallback_models[name] = model
181192
return model
@@ -258,7 +269,11 @@ def _process_fallback(model_name, audio_path, language=None):
258269
return {"text": text, "segments": segments, "duration": duration}
259270

260271
# Pre-load
261-
if preload_model == "fun-asr-nano":
272+
if app.state.model_path:
273+
# When custom model_path is provided, use it as the model name for loading
274+
logger.info(f"Loading custom model: {app.state.model_path} (hub: {app.state.hub})")
275+
_load_fallback("custom")
276+
elif preload_model == "fun-asr-nano":
262277
_load_vllm_engine()
263278
else:
264279
_load_fallback(preload_model)
@@ -288,7 +303,7 @@ async def transcribe(
288303
result = _process_fallback("fun-asr-nano", tmp_path, language=language)
289304
finally:
290305
os.unlink(tmp_path)
291-
elif model in FALLBACK_CONFIGS:
306+
elif model in FALLBACK_CONFIGS or model == "custom":
292307
suffix = os.path.splitext(file.filename)[1] if file.filename else ".wav"
293308
with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp:
294309
tmp.write(content)
@@ -298,7 +313,8 @@ async def transcribe(
298313
finally:
299314
os.unlink(tmp_path)
300315
else:
301-
raise HTTPException(400, f"Unknown model '{model}'. Available: fun-asr-nano, {', '.join(FALLBACK_CONFIGS.keys())}")
316+
available = ["fun-asr-nano", "custom"] + list(FALLBACK_CONFIGS.keys())
317+
raise HTTPException(400, f"Unknown model '{model}'. Available: {', '.join(available)}")
302318

303319
t1 = time.perf_counter()
304320

@@ -352,6 +368,8 @@ async def asr_endpoint(
352368
@app.get("/v1/models")
353369
async def list_models():
354370
all_models = ["fun-asr-nano"] + list(FALLBACK_CONFIGS.keys())
371+
if app.state.model_path:
372+
all_models.append("custom")
355373
return JSONResponse({"object": "list", "data": [{"id": n, "object": "model"} for n in all_models]})
356374

357375
@app.get("/health")

funasr/bin/server.py

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,8 @@
55
funasr-server # default: sensevoice on cuda:0, port 8000
66
funasr-server --device cpu --port 9000
77
funasr-server --model paraformer
8+
funasr-server --model-path /path/to/local/model
9+
funasr-server --model-path username/paraformer --hub hf
810
"""
911

1012
import argparse
@@ -21,6 +23,8 @@ def main():
2123
funasr-server --device cpu # Start on CPU
2224
funasr-server --model paraformer # Use Paraformer model
2325
funasr-server --port 9000 # Custom port
26+
funasr-server --model-path /path/to/local/model # Use local model
27+
funasr-server --model-path username/model --hub hf # Use HuggingFace model
2428
2529
Then use with OpenAI SDK:
2630
from openai import OpenAI
@@ -32,6 +36,8 @@ def main():
3236
parser.add_argument("--port", type=int, default=8000, help="Port (default: 8000)")
3337
parser.add_argument("--device", default="cuda", help="Device: cuda, cpu, mps (default: cuda)")
3438
parser.add_argument("--model", default="auto", help="Pre-load model: auto (GPU=fun-asr-nano, CPU=sensevoice), sensevoice, paraformer, fun-asr-nano")
39+
parser.add_argument("--model-path", default=None, help="Local model path or model ID (overrides --model)")
40+
parser.add_argument("--hub", default="ms", help="Model hub: ms (ModelScope), hf (HuggingFace) (default: ms)")
3541
args = parser.parse_args()
3642

3743
try:
@@ -49,12 +55,15 @@ def main():
4955
# Use inline app to avoid path issues
5056
from funasr.bin._server_app import create_app
5157

52-
app = create_app(device=args.device, preload_model=args.model)
58+
app = create_app(device=args.device, preload_model=args.model, model_path=args.model_path, hub=args.hub)
5359

5460
print(f"╔══════════════════════════════════════════════╗")
5561
print(f"║ FunASR Server v1.3.6 ║")
5662
print(f"║ Device: {args.device:<8} ║")
5763
print(f"║ Model: {args.model:<12} ║")
64+
if args.model_path:
65+
print(f"║ Model Path: {args.model_path:<25} ║")
66+
print(f"║ Hub: {args.hub:<8} ║")
5867
print(f"║ URL: http://{args.host}:{args.port}/v1 ║")
5968
print(f"║ Docs: http://{args.host}:{args.port}/docs ║")
6069
print(f"╚══════════════════════════════════════════════╝")

tests/test_server_app_openai_segments.py

Lines changed: 72 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,8 +9,18 @@
99

1010

1111
def load_server_app(monkeypatch):
12+
class DummyFastAPI:
13+
def __init__(self, *args, **kwargs):
14+
self.state = types.SimpleNamespace()
15+
16+
def post(self, *args, **kwargs):
17+
return lambda func: func
18+
19+
def get(self, *args, **kwargs):
20+
return lambda func: func
21+
1222
fastapi_stub = types.ModuleType("fastapi")
13-
fastapi_stub.FastAPI = object
23+
fastapi_stub.FastAPI = DummyFastAPI
1424
fastapi_stub.UploadFile = object
1525
fastapi_stub.File = lambda *args, **kwargs: None
1626
fastapi_stub.Form = lambda *args, **kwargs: None
@@ -31,6 +41,39 @@ def load_server_app(monkeypatch):
3141
return module
3242

3343

44+
def install_dummy_funasr(monkeypatch):
45+
class DummyAutoModel:
46+
instances = []
47+
48+
def __init__(self, **kwargs):
49+
self.kwargs = kwargs
50+
self.__class__.instances.append(kwargs)
51+
52+
funasr_stub = types.ModuleType("funasr")
53+
funasr_stub.AutoModel = DummyAutoModel
54+
monkeypatch.setitem(sys.modules, "funasr", funasr_stub)
55+
return DummyAutoModel
56+
57+
58+
def install_dummy_vllm(monkeypatch):
59+
class DummyVLLM:
60+
calls = []
61+
62+
@classmethod
63+
def from_pretrained(cls, **kwargs):
64+
cls.calls.append(kwargs)
65+
return object()
66+
67+
monkeypatch.setitem(sys.modules, "funasr.models", types.ModuleType("funasr.models"))
68+
monkeypatch.setitem(
69+
sys.modules, "funasr.models.fun_asr_nano", types.ModuleType("funasr.models.fun_asr_nano")
70+
)
71+
vllm_module = types.ModuleType("funasr.models.fun_asr_nano.inference_vllm")
72+
vllm_module.FunASRNanoVLLM = DummyVLLM
73+
monkeypatch.setitem(sys.modules, "funasr.models.fun_asr_nano.inference_vllm", vllm_module)
74+
return DummyVLLM
75+
76+
3477
def test_fallback_segments_split_long_fun_asr_server_text(monkeypatch):
3578
module = load_server_app(monkeypatch)
3679
text = (
@@ -56,3 +99,31 @@ def test_fallback_segments_keep_short_text_single_cue(monkeypatch):
5699
assert module.build_openai_fallback_segments("hello", duration=1.25) == [
57100
{"start": 0.0, "end": 1.25, "text": "hello"}
58101
]
102+
103+
104+
def test_default_fun_asr_nano_uses_huggingface_hub(monkeypatch):
105+
module = load_server_app(monkeypatch)
106+
DummyAutoModel = install_dummy_funasr(monkeypatch)
107+
DummyVLLM = install_dummy_vllm(monkeypatch)
108+
109+
module.create_app(device="cuda", preload_model="fun-asr-nano", hub="ms")
110+
111+
assert DummyVLLM.calls[0]["model"] == "FunAudioLLM/Fun-ASR-Nano-2512"
112+
assert DummyVLLM.calls[0]["hub"] == "hf"
113+
assert DummyAutoModel.instances[0]["model"] == "fsmn-vad"
114+
115+
116+
def test_custom_model_path_fallback_uses_empty_config_and_requested_hub(monkeypatch):
117+
module = load_server_app(monkeypatch)
118+
DummyAutoModel = install_dummy_funasr(monkeypatch)
119+
120+
app = module.create_app(
121+
device="cpu",
122+
preload_model="sensevoice",
123+
model_path="org/custom-sensevoice",
124+
hub="hf",
125+
)
126+
127+
assert DummyAutoModel.instances[0]["model"] == "org/custom-sensevoice"
128+
assert DummyAutoModel.instances[0]["hub"] == "hf"
129+
assert app.state.fallback_models["custom"] is not None

0 commit comments

Comments
 (0)