@@ -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" )
0 commit comments