| 35 | sys.path.extend([main_root, third_party_matcha_tts_path]) |
| 36 | |
| 37 | class InspireMusicModel: |
| 38 | def __init__(self, |
| 39 | model_name: str, |
| 40 | model_dir: str = None, |
| 41 | min_generate_audio_seconds: float = 0.0, |
| 42 | max_generate_audio_seconds: float = 30.0, |
| 43 | sample_rate: int = 24000, |
| 44 | output_sample_rate: int = 48000, |
| 45 | load_jit: bool = True, |
| 46 | load_onnx: bool = False, |
| 47 | dtype: str = "fp16", |
| 48 | fast: bool = False, |
| 49 | fp16: bool = True, |
| 50 | gpu: int = 1, |
| 51 | result_dir: str = None, |
| 52 | hub="modelscope", |
| 53 | repo_url=None, |
| 54 | token=None): |
| 55 | os.environ['CUDA_VISIBLE_DEVICES'] = str(gpu) |
| 56 | |
| 57 | # Set model_dir or default to downloading if it doesn't exist |
| 58 | if model_dir is None: |
| 59 | if sys.platform == "win32": |
| 60 | model_dir = f"..\..\pretrained_models\{model_name}" |
| 61 | else: |
| 62 | model_dir = f"../../pretrained_models/{model_name}" |
| 63 | |
| 64 | if not os.path.isfile(os.path.join(model_dir, "llm.pt")): |
| 65 | if hub == "modelscope": |
| 66 | from modelscope import snapshot_download |
| 67 | if model_name == "InspireMusic-Base": |
| 68 | snapshot_download(f"iic/InspireMusic", local_dir=model_dir) |
| 69 | else: |
| 70 | snapshot_download(f"iic/{model_name}", local_dir=model_dir) |
| 71 | elif hub == "huggingface": |
| 72 | from huggingface_hub import snapshot_download |
| 73 | snapshot_download(repo_id=f"FunAudioLLM/{model_name}", local_dir=model_dir) |
| 74 | |
| 75 | self.model_dir = model_dir |
| 76 | |
| 77 | self.sample_rate = sample_rate |
| 78 | self.output_sample_rate = 24000 if fast else output_sample_rate |
| 79 | self.result_dir = result_dir or os.path.join("exp", model_name) |
| 80 | os.makedirs(self.result_dir, exist_ok=True) |
| 81 | |
| 82 | self.min_generate_audio_seconds = min_generate_audio_seconds |
| 83 | self.max_generate_audio_seconds = max_generate_audio_seconds |
| 84 | self.min_generate_audio_length = int(self.output_sample_rate * self.min_generate_audio_seconds) |
| 85 | self.max_generate_audio_length = int(self.output_sample_rate * self.max_generate_audio_seconds) |
| 86 | assert self.min_generate_audio_seconds <= self.max_generate_audio_seconds, "Min audio seconds must be less than or equal to max audio seconds" |
| 87 | |
| 88 | use_cuda = gpu >= 0 and torch.cuda.is_available() |
| 89 | if gpu >=0: |
| 90 | if torch.cuda.is_available(): |
| 91 | self.device = torch.device('cuda') |
| 92 | elif torch.backends.mps.is_available(): |
| 93 | self.device = torch.device('mps') |
| 94 | elif torch.xpu.is_available(): |
no outgoing calls
no test coverage detected