(self,
llm: torch.nn.Module,
flow: torch.nn.Module,
music_tokenizer: torch.nn.Module,
wavtokenizer: torch.nn.Module,
dtype: str = "fp16",
fast: bool = False,
fp16: bool = True,
)
| 28 | |
| 29 | class InspireMusicModel: |
| 30 | def __init__(self, |
| 31 | llm: torch.nn.Module, |
| 32 | flow: torch.nn.Module, |
| 33 | music_tokenizer: torch.nn.Module, |
| 34 | wavtokenizer: torch.nn.Module, |
| 35 | dtype: str = "fp16", |
| 36 | fast: bool = False, |
| 37 | fp16: bool = True, |
| 38 | ): |
| 39 | |
| 40 | if torch.cuda.is_available(): |
| 41 | self.device = torch.device('cuda') |
| 42 | elif torch.backends.mps.is_available(): |
| 43 | self.device = torch.device('mps') |
| 44 | elif torch.xpu.is_available(): |
| 45 | self.device = torch.device('xpu') |
| 46 | else: |
| 47 | self.device = torch.device('cpu') |
| 48 | |
| 49 | if dtype == "fp16": |
| 50 | self.dtype = torch.float16 |
| 51 | elif dtype == "bf16": |
| 52 | self.dtype = torch.bfloat16 |
| 53 | else: |
| 54 | self.dtype = torch.float32 |
| 55 | |
| 56 | self.llm = llm.to(self.dtype) |
| 57 | self.flow = flow |
| 58 | self.music_tokenizer = music_tokenizer |
| 59 | self.wavtokenizer = wavtokenizer |
| 60 | self.fp16 = fp16 |
| 61 | self.token_min_hop_len = 100 |
| 62 | self.token_max_hop_len = 200 |
| 63 | self.token_overlap_len = 20 |
| 64 | # mel fade in out |
| 65 | self.mel_overlap_len = 34 |
| 66 | self.mel_window = np.hamming(2 * self.mel_overlap_len) |
| 67 | # hift cache |
| 68 | self.mel_cache_len = 20 |
| 69 | self.source_cache_len = int(self.mel_cache_len * 256) |
| 70 | # rtf and decoding related |
| 71 | self.stream_scale_factor = 1 |
| 72 | assert self.stream_scale_factor >= 1, 'stream_scale_factor should be greater than 1, change it according to your actual rtf' |
| 73 | self.llm_context = torch.cuda.stream(torch.cuda.Stream(self.device)) if torch.cuda.is_available() else nullcontext() |
| 74 | self.lock = threading.Lock() |
| 75 | # dict used to store session related variable |
| 76 | self.music_token_dict = {} |
| 77 | self.llm_end_dict = {} |
| 78 | self.mel_overlap_dict = {} |
| 79 | self.fast = fast |
| 80 | self.generator = "hifi" |
| 81 | |
| 82 | def load(self, llm_model, flow_model, hift_model, wavtokenizer_model): |
| 83 | if llm_model is not None: |
nothing calls this directly
no outgoing calls
no test coverage detected