MCPcopy Create free account
hub / github.com/FunAudioLLM/FunMusic / InspireMusicModel

Class InspireMusicModel

inspiremusic/cli/inference.py:37–220  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

35 sys.path.extend([main_root, third_party_matcha_tts_path])
36
37class 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():

Callers 2

music_generationFunction · 0.90
mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected