2222
2323import subprocess
2424
25- from . import audio_legacy
25+ from . import audio_legacy # noqa: F401
2626import torch as th
2727import torchaudio as ta
2828
@@ -118,9 +118,17 @@ def __init__(
118118 self ._name = model
119119 self ._repo = repo
120120 self ._load_model ()
121- self .update_parameter (device = device , shifts = shifts , overlap = overlap , split = split ,
122- segment = segment , jobs = jobs , progress = progress , callback = callback ,
123- callback_arg = callback_arg )
121+ self .update_parameter (
122+ device = device ,
123+ shifts = shifts ,
124+ overlap = overlap ,
125+ split = split ,
126+ segment = segment ,
127+ jobs = jobs ,
128+ progress = progress ,
129+ callback = callback ,
130+ callback_arg = callback_arg ,
131+ )
124132
125133 def update_parameter (
126134 self ,
@@ -131,9 +139,7 @@ def update_parameter(
131139 segment : Optional [Union [int , _NotProvided ]] = NotProvided ,
132140 jobs : Union [int , _NotProvided ] = NotProvided ,
133141 progress : Union [bool , _NotProvided ] = NotProvided ,
134- callback : Optional [
135- Union [Callable [[dict ], None ], _NotProvided ]
136- ] = NotProvided ,
142+ callback : Optional [Union [Callable [[dict ], None ], _NotProvided ]] = NotProvided ,
137143 callback_arg : Optional [Union [dict , _NotProvided ]] = NotProvided ,
138144 ):
139145 """
@@ -213,8 +219,9 @@ def _load_audio(self, track: Path):
213219 wav = None
214220
215221 try :
216- wav = AudioFile (track ).read (streams = 0 , samplerate = self ._samplerate ,
217- channels = self ._audio_channels )
222+ wav = AudioFile (track ).read (
223+ streams = 0 , samplerate = self ._samplerate , channels = self ._audio_channels
224+ )
218225 except FileNotFoundError :
219226 errors ["ffmpeg" ] = "FFmpeg is not installed."
220227 except subprocess .CalledProcessError :
@@ -269,20 +276,20 @@ def separate_tensor(
269276 wav -= ref .mean ()
270277 wav /= ref .std () + 1e-8
271278 out = apply_model (
272- self ._model ,
273- wav [None ],
274- segment = self ._segment ,
275- shifts = self ._shifts ,
276- split = self ._split ,
277- overlap = self ._overlap ,
278- device = self ._device ,
279- num_workers = self ._jobs ,
280- callback = self ._callback ,
281- callback_arg = _replace_dict (
282- self ._callback_arg , ("audio_length" , wav .shape [1 ])
283- ),
284- progress = self ._progress ,
285- )
279+ self ._model ,
280+ wav [None ],
281+ segment = self ._segment ,
282+ shifts = self ._shifts ,
283+ split = self ._split ,
284+ overlap = self ._overlap ,
285+ device = self ._device ,
286+ num_workers = self ._jobs ,
287+ callback = self ._callback ,
288+ callback_arg = _replace_dict (
289+ self ._callback_arg , ("audio_length" , wav .shape [1 ])
290+ ),
291+ progress = self ._progress ,
292+ )
286293 if out is None :
287294 raise KeyboardInterrupt
288295 out *= ref .std () + 1e-8
@@ -336,7 +343,7 @@ def list_models(repo: Optional[Path] = None) -> Dict[str, Dict[str, Union[str, P
336343 """
337344 model_repo : ModelOnlyRepo
338345 if repo is None :
339- models = _parse_remote_files (REMOTE_ROOT / ' files.txt' )
346+ models = _parse_remote_files (REMOTE_ROOT / " files.txt" )
340347 model_repo = RemoteRepo (models )
341348 bag_repo = BagOnlyRepo (REMOTE_ROOT , model_repo )
342349 else :
@@ -363,7 +370,7 @@ def list_models(repo: Optional[Path] = None) -> Dict[str, Dict[str, Union[str, P
363370 split = args .split ,
364371 segment = args .segment ,
365372 jobs = args .jobs ,
366- callback = print
373+ callback = print ,
367374 )
368375 out = args .out / args .name
369376 out .mkdir (parents = True , exist_ok = True )
0 commit comments