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

Method __init__

inspiremusic/cli/inference.py:38–101  ·  view source on GitHub ↗
(self,
                 model_name: str,
                 model_dir: str = None,
                 min_generate_audio_seconds: float = 0.0,
                 max_generate_audio_seconds: float = 30.0,
                 sample_rate: int = 24000,
                 output_sample_rate: int = 48000,
                 load_jit: bool = True,
                 load_onnx: bool = False,
                 dtype: str = "fp16",
                 fast: bool = False,
                 fp16: bool = True,
                 gpu: int = 1,
                 result_dir: str = None,
                 hub="modelscope",
                 repo_url=None,
                 token=None)

Source from the content-addressed store, hash-verified

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():
95 self.device = torch.device('xpu')

Callers

nothing calls this directly

Calls 1

InspireMusicClass · 0.90

Tested by

no test coverage detected