Instructions to use SimpleTuner/open-rvq-encoder-minimax-music3 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use SimpleTuner/open-rvq-encoder-minimax-music3 with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("SimpleTuner/open-rvq-encoder-minimax-music3", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Keep only periodic c0 reference constraints
Browse files- README.md +20 -17
- comfyui_open_rvq/nodes.py +109 -46
- comfyui_workflow_example.json +11 -2
- minimax_music3_reference_adapter.py +171 -52
README.md
CHANGED
|
@@ -27,21 +27,16 @@ The project started with a 41M-parameter, single-GPU community proof at **0.6633
|
|
| 27 |
|
| 28 |
Use **v4** unless reproducing an experiment.
|
| 29 |
|
|
|
|
|
|
|
| 30 |
## Files
|
| 31 |
|
| 32 |
-
|
|
| 33 |
|---|---:|---|---:|
|
| 34 |
-
|
|
| 35 |
-
|
|
| 36 |
-
|
|
| 37 |
-
|
|
| 38 |
-
|
| 39 |
-
Weight files in [`encoders/`](encoders/):
|
| 40 |
-
|
| 41 |
-
- **v1**: `minimax_music3_rvq_encoder_v1_41m_independent_heads.safetensors`
|
| 42 |
-
- **v2**: `minimax_music3_rvq_encoder_v2_155m_wide_independent_heads.safetensors`
|
| 43 |
-
- **v3**: `minimax_music3_rvq_encoder_v3_155m_mert_aligned_independent_heads.safetensors`
|
| 44 |
-
- **v4**: `minimax_music3_rvq_encoder_v4_169m_autoregressive_depth_recommended.safetensors`
|
| 45 |
|
| 46 |
Each weight file has a same-named `.json` configuration file in [`encoders/`](encoders/).
|
| 47 |
|
|
@@ -136,7 +131,14 @@ ln -s /path/to/open-rvq-encoder-minimax-music3/comfyui_open_rvq \
|
|
| 136 |
custom_nodes/comfyui_open_rvq
|
| 137 |
```
|
| 138 |
|
| 139 |
-
Restart ComfyUI.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 140 |
|
| 141 |
The node package reads the encoder files directly from this clone. They can instead be placed in:
|
| 142 |
|
|
@@ -144,7 +146,7 @@ The node package reads the encoder files directly from this clone. They can inst
|
|
| 144 |
ComfyUI/models/minimax_music3_rvq_encoders/
|
| 145 |
```
|
| 146 |
|
| 147 |
-
The
|
| 148 |
|
| 149 |
## Diffusers
|
| 150 |
|
|
@@ -190,6 +192,7 @@ frame_hiddens, predicted_codes = adapter.encode_reference(
|
|
| 190 |
lyrics="[instrumental]",
|
| 191 |
generator=generator,
|
| 192 |
device="cuda",
|
|
|
|
| 193 |
)
|
| 194 |
|
| 195 |
result = pipe(
|
|
@@ -200,7 +203,7 @@ result = pipe(
|
|
| 200 |
)
|
| 201 |
```
|
| 202 |
|
| 203 |
-
The patch only adds a precomputed-`frame_hiddens` bypass to the modular pipeline. It does not replace MiniMax model code.
|
| 204 |
|
| 205 |
## Limits
|
| 206 |
|
|
@@ -209,8 +212,8 @@ The patch only adds a precomputed-`frame_hiddens` bypass to the modular pipeline
|
|
| 209 |
- Training data is synthetic MiniMax Music 3 output, not MiniMax's training set.
|
| 210 |
- Real-audio generalization is not established.
|
| 211 |
- Context is 128 frames, or 5.12 seconds. There is no cross-window encoder state.
|
| 212 |
-
- Reference
|
| 213 |
-
-
|
| 214 |
|
| 215 |
## Credits
|
| 216 |
|
|
|
|
| 27 |
|
| 28 |
Use **v4** unless reproducing an experiment.
|
| 29 |
|
| 30 |
+
The released ComfyUI and Diffusers adapters support one reference-generation method. Every fifth generated semantic `c0` token is restricted to the encoder's top-5 candidates by default. MiniMax chooses the token and generates all acoustic codebooks. The interval is configurable from 1 through 10.
|
| 31 |
+
|
| 32 |
## Files
|
| 33 |
|
| 34 |
+
| File | Parameters | Experiment | Replay cosine |
|
| 35 |
|---|---:|---|---:|
|
| 36 |
+
| `minimax_music3_rvq_encoder_v1_41m_independent_heads.safetensors` | 40,978,944 | Baseline; eight independent heads | 0.7624* |
|
| 37 |
+
| `minimax_music3_rvq_encoder_v2_155m_wide_independent_heads.safetensors` | 154,736,064 | Wider shared encoder | 0.7698 |
|
| 38 |
+
| `minimax_music3_rvq_encoder_v3_155m_mert_aligned_independent_heads.safetensors` | 154,736,064 | v2 plus training-only MERT alignment | 0.7703 |
|
| 39 |
+
| `minimax_music3_rvq_encoder_v4_169m_autoregressive_depth_recommended.safetensors` | 169,008,576 | Causal acoustic decoder across codebook depth | **0.8748** |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 40 |
|
| 41 |
Each weight file has a same-named `.json` configuration file in [`encoders/`](encoders/).
|
| 42 |
|
|
|
|
| 131 |
custom_nodes/comfyui_open_rvq
|
| 132 |
```
|
| 133 |
|
| 134 |
+
Restart ComfyUI. Upload a reference audio file and select v4 in **MiniMax Music3 RVQ Reference Encoder Loader**.
|
| 135 |
+
|
| 136 |
+
- [`comfyui_workflow_example.json`](comfyui_workflow_example.json) constrains every fifth generated semantic `c0` token to the encoder's top-5 candidates.
|
| 137 |
+
- `reference_interval=1` constrains every frame.
|
| 138 |
+
- `reference_interval=5` is the tested default.
|
| 139 |
+
- `reference_interval=10` constrains every tenth frame and gives the language model more freedom.
|
| 140 |
+
|
| 141 |
+
The seven acoustic codebooks are generated by MiniMax. They are not copied from the reference. Describe the target arrangement in `caption`. Provide the desired sectioned lyrics. Prompt adherence and audio quality vary; this is not a general audio-to-audio conversion system.
|
| 142 |
|
| 143 |
The node package reads the encoder files directly from this clone. They can instead be placed in:
|
| 144 |
|
|
|
|
| 146 |
ComfyUI/models/minimax_music3_rvq_encoders/
|
| 147 |
```
|
| 148 |
|
| 149 |
+
The interval-5 graph was verified with a 30-second reference, five Euler steps, the pruned int8 text encoder, the fp16 diffusion model, and the v4 encoder. It produced a 29.99-second stereo 44.1 kHz FLAC. Use 30 diffusion steps for normal output.
|
| 150 |
|
| 151 |
## Diffusers
|
| 152 |
|
|
|
|
| 192 |
lyrics="[instrumental]",
|
| 193 |
generator=generator,
|
| 194 |
device="cuda",
|
| 195 |
+
reference_interval=5,
|
| 196 |
)
|
| 197 |
|
| 198 |
result = pipe(
|
|
|
|
| 203 |
)
|
| 204 |
```
|
| 205 |
|
| 206 |
+
`reference_interval` accepts integers from 1 through 10. The patch only adds a precomputed-`frame_hiddens` bypass to the modular pipeline. It does not replace MiniMax model code.
|
| 207 |
|
| 208 |
## Limits
|
| 209 |
|
|
|
|
| 212 |
- Training data is synthetic MiniMax Music 3 output, not MiniMax's training set.
|
| 213 |
- Real-audio generalization is not established.
|
| 214 |
- Context is 128 frames, or 5.12 seconds. There is no cross-window encoder state.
|
| 215 |
+
- Reference generation still needs the official MiniMax language model and RVQ depth decoder.
|
| 216 |
+
- The released integration uses the RVQ encoder's top-5 semantic candidates. Other candidate counts are not exposed.
|
| 217 |
|
| 218 |
## Credits
|
| 219 |
|
comfyui_open_rvq/nodes.py
CHANGED
|
@@ -15,9 +15,9 @@ import torch
|
|
| 15 |
|
| 16 |
import comfy.model_management
|
| 17 |
import comfy.model_prefetch
|
| 18 |
-
import comfy.ops
|
| 19 |
import folder_paths
|
| 20 |
-
from comfy.ldm.minimax_music.ar import AUDIO_FRAMES_PER_SECOND, CFG_SCALE, CFG_TOP_K, MiniMaxMusic3AR
|
|
|
|
| 21 |
from comfy.text_encoders.minimax_music import MiniMaxMusic3TEModel
|
| 22 |
|
| 23 |
|
|
@@ -29,30 +29,18 @@ from minimax_music3_reference_adapter import MiniMaxMusic3ReferenceAdapter # no
|
|
| 29 |
|
| 30 |
|
| 31 |
ENCODER_FOLDER = "minimax_music3_rvq_encoders"
|
| 32 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 33 |
bundled_encoders = REPO_ROOT / "encoders"
|
| 34 |
if bundled_encoders.is_dir():
|
| 35 |
folder_paths.add_model_folder_path(ENCODER_FOLDER, str(bundled_encoders))
|
| 36 |
|
| 37 |
|
| 38 |
-
def _teacher_depth_hidden(model: MiniMaxMusic3AR, hidden, codes, execution_dtype):
|
| 39 |
-
decoder = model.model.audio_decoder
|
| 40 |
-
sequence = [decoder.projection(hidden).unsqueeze(1)]
|
| 41 |
-
semantic = model._embed_c0(codes[:, 0], execution_dtype)
|
| 42 |
-
sequence.append(decoder.projection(semantic).unsqueeze(1))
|
| 43 |
-
hidden_parts = []
|
| 44 |
-
for index in range(1, model.num_codebooks):
|
| 45 |
-
depth_hidden = decoder(torch.cat(sequence, dim=1))[:, -1]
|
| 46 |
-
hidden_parts.append(depth_hidden[:1])
|
| 47 |
-
if index < model.num_codebooks - 1:
|
| 48 |
-
embedding = model.model.audio_extra_embedding(
|
| 49 |
-
codes[:, index] + (index - 1) * model.audio_vocab_size,
|
| 50 |
-
out_dtype=execution_dtype,
|
| 51 |
-
)
|
| 52 |
-
sequence.append(decoder.projection(embedding).unsqueeze(1))
|
| 53 |
-
return torch.cat(hidden_parts, dim=-1)
|
| 54 |
-
|
| 55 |
-
|
| 56 |
def _run_depth(model: MiniMaxMusic3AR, device, execution_dtype, core):
|
| 57 |
decoder = model.model.audio_decoder
|
| 58 |
queue = comfy.model_prefetch.make_prefetch_queue(
|
|
@@ -71,25 +59,59 @@ def _run_depth(model: MiniMaxMusic3AR, device, execution_dtype, core):
|
|
| 71 |
comfy.model_prefetch.prefetch_queue_pop(queue, device, None)
|
| 72 |
|
| 73 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 74 |
@torch.inference_mode()
|
| 75 |
-
def replay_reference_codes(
|
| 76 |
-
|
| 77 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 78 |
prompt_tokens = int(input_ids.shape[1])
|
| 79 |
input_ids = input_ids.to(device)
|
| 80 |
-
|
| 81 |
execution_dtype = torch.bfloat16 if comfy.model_management.should_use_bf16(device) else torch.float32
|
| 82 |
unconditioned = input_ids.clone()
|
| 83 |
-
from comfy.ldm.minimax_music.prompt import
|
| 84 |
|
| 85 |
unconditioned[:, 1:-2] = SPECIAL_TOKEN_IDS["<|audio_cfg|>"]
|
| 86 |
text_ids = torch.cat((input_ids, unconditioned), dim=0)
|
| 87 |
-
|
| 88 |
-
|
| 89 |
-
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
|
|
|
|
|
|
|
| 93 |
last_hidden, past = output[0][:, -1], output[2]
|
| 94 |
from comfy.ldm.minimax_music.ar import derive_seed
|
| 95 |
|
|
@@ -123,19 +145,41 @@ def replay_reference_codes(model: MiniMaxMusic3AR, input_ids, codes, seed, devic
|
|
| 123 |
output = model.model(None, embeds=feedback, past_key_values=past, dtype=execution_dtype)
|
| 124 |
last_hidden, past = output[0][:, -1], output[2]
|
| 125 |
|
|
|
|
| 126 |
frames = []
|
| 127 |
-
for frame_index in range(
|
| 128 |
comfy.model_management.throw_exception_if_processing_interrupted()
|
| 129 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 130 |
result: dict[str, torch.Tensor] = {}
|
| 131 |
|
| 132 |
-
def
|
| 133 |
-
result["hidden"] =
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 134 |
|
| 135 |
-
_run_depth(model, device, execution_dtype,
|
| 136 |
frames.append(torch.cat((last_hidden[:1], result["hidden"]), dim=-1)[0].cpu())
|
| 137 |
-
if frame_index + 1 <
|
| 138 |
-
feedback = model._embed_audio_frame(
|
| 139 |
output = model.model(None, embeds=feedback, past_key_values=past, dtype=execution_dtype)
|
| 140 |
last_hidden, past = output[0][:, -1], output[2]
|
| 141 |
return torch.stack(frames)
|
|
@@ -145,19 +189,20 @@ _ORIGINAL_ENCODE_TOKEN_WEIGHTS = MiniMaxMusic3TEModel.encode_token_weights
|
|
| 145 |
|
| 146 |
|
| 147 |
def _encode_token_weights_with_reference(self, token_weight_pairs):
|
| 148 |
-
|
| 149 |
-
if
|
| 150 |
return _ORIGINAL_ENCODE_TOKEN_WEIGHTS(self, token_weight_pairs)
|
| 151 |
token_ids = [token for token, _ in token_weight_pairs["minimax_music3"][0]]
|
| 152 |
input_ids = torch.tensor([token_ids], dtype=torch.long)
|
| 153 |
hidden = replay_reference_codes(
|
| 154 |
self,
|
| 155 |
input_ids,
|
| 156 |
-
|
| 157 |
int(token_weight_pairs["seed"]),
|
| 158 |
self.execution_device,
|
| 159 |
float(token_weight_pairs["cfg_scale"]),
|
| 160 |
int(token_weight_pairs["top_k"]),
|
|
|
|
| 161 |
)
|
| 162 |
return hidden.unsqueeze(0), None, {}
|
| 163 |
|
|
@@ -170,7 +215,11 @@ if not getattr(MiniMaxMusic3TEModel, "_simpletuner_reference_patch", False):
|
|
| 170 |
class MiniMaxMusic3RVQReferenceEncoderLoader:
|
| 171 |
@classmethod
|
| 172 |
def INPUT_TYPES(cls):
|
| 173 |
-
encoders = [
|
|
|
|
|
|
|
|
|
|
|
|
|
| 174 |
dav_files = [name for name in folder_paths.get_filename_list("vae") if name.lower().endswith((".pth", ".pt"))]
|
| 175 |
return {"required": {"encoder": (encoders,), "dav": (dav_files,)}}
|
| 176 |
|
|
@@ -199,6 +248,7 @@ class MiniMaxMusic3ReferenceAudioEncode:
|
|
| 199 |
"caption": ("STRING", {"multiline": True, "dynamicPrompts": True}),
|
| 200 |
"lyrics": ("STRING", {"multiline": True, "dynamicPrompts": True}),
|
| 201 |
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
|
|
|
|
| 202 |
},
|
| 203 |
"optional": {
|
| 204 |
"cfg_scale": ("FLOAT", {"default": CFG_SCALE, "min": 0.0, "max": 100.0, "step": 0.1}),
|
|
@@ -211,10 +261,22 @@ class MiniMaxMusic3ReferenceAudioEncode:
|
|
| 211 |
FUNCTION = "encode"
|
| 212 |
CATEGORY = "conditioning/minimax music"
|
| 213 |
|
| 214 |
-
def encode(
|
| 215 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 216 |
audio["waveform"],
|
| 217 |
int(audio["sample_rate"]),
|
|
|
|
| 218 |
device=comfy.model_management.get_torch_device(),
|
| 219 |
)
|
| 220 |
tokens = clip.tokenize(
|
|
@@ -225,7 +287,8 @@ class MiniMaxMusic3ReferenceAudioEncode:
|
|
| 225 |
cfg_scale=cfg_scale,
|
| 226 |
top_k=top_k,
|
| 227 |
)
|
| 228 |
-
tokens["
|
|
|
|
| 229 |
conditioning = clip.encode_from_tokens_scheduled(tokens)
|
| 230 |
for cond in conditioning:
|
| 231 |
hidden = cond[0]
|
|
|
|
| 15 |
|
| 16 |
import comfy.model_management
|
| 17 |
import comfy.model_prefetch
|
|
|
|
| 18 |
import folder_paths
|
| 19 |
+
from comfy.ldm.minimax_music.ar import AUDIO_FRAMES_PER_SECOND, CFG_SCALE, CFG_TOP_K, MiniMaxMusic3AR, sample_topk
|
| 20 |
+
from comfy.ldm.minimax_music.prompt import AUDIO_CODE_OFFSET
|
| 21 |
from comfy.text_encoders.minimax_music import MiniMaxMusic3TEModel
|
| 22 |
|
| 23 |
|
|
|
|
| 29 |
|
| 30 |
|
| 31 |
ENCODER_FOLDER = "minimax_music3_rvq_encoders"
|
| 32 |
+
REFERENCE_CANDIDATE_COUNT = 5
|
| 33 |
+
DEFAULT_REFERENCE_INTERVAL = 5
|
| 34 |
+
folder_paths.add_model_folder_path(
|
| 35 |
+
ENCODER_FOLDER,
|
| 36 |
+
str(Path(folder_paths.models_dir) / ENCODER_FOLDER),
|
| 37 |
+
is_default=True,
|
| 38 |
+
)
|
| 39 |
bundled_encoders = REPO_ROOT / "encoders"
|
| 40 |
if bundled_encoders.is_dir():
|
| 41 |
folder_paths.add_model_folder_path(ENCODER_FOLDER, str(bundled_encoders))
|
| 42 |
|
| 43 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 44 |
def _run_depth(model: MiniMaxMusic3AR, device, execution_dtype, core):
|
| 45 |
decoder = model.model.audio_decoder
|
| 46 |
queue = comfy.model_prefetch.make_prefetch_queue(
|
|
|
|
| 59 |
comfy.model_prefetch.prefetch_queue_pop(queue, device, None)
|
| 60 |
|
| 61 |
|
| 62 |
+
def _sample_semantic_candidates(model, hidden, semantic_candidates, cfg_scale, top_k, generator):
|
| 63 |
+
if model.model.pruned_lm_head:
|
| 64 |
+
logits = model.model.lm_head_pruned(hidden).float()
|
| 65 |
+
token_candidates = semantic_candidates + 1
|
| 66 |
+
offset = 1
|
| 67 |
+
else:
|
| 68 |
+
logits = model.model.lm_head(hidden).float()
|
| 69 |
+
token_candidates = semantic_candidates + AUDIO_CODE_OFFSET
|
| 70 |
+
offset = AUDIO_CODE_OFFSET
|
| 71 |
+
token_candidates = token_candidates.to(device=logits.device, dtype=torch.long).unsqueeze(0)
|
| 72 |
+
conditioned = logits[:1].gather(-1, token_candidates)
|
| 73 |
+
unconditioned = logits[1:2].gather(-1, token_candidates)
|
| 74 |
+
guided = unconditioned + (conditioned - unconditioned) * cfg_scale
|
| 75 |
+
selected_index = sample_topk(guided, min(top_k, guided.shape[-1]), generator).view(1, 1)
|
| 76 |
+
return token_candidates.gather(-1, selected_index).squeeze(-1) - offset
|
| 77 |
+
|
| 78 |
+
|
| 79 |
@torch.inference_mode()
|
| 80 |
+
def replay_reference_codes(
|
| 81 |
+
model,
|
| 82 |
+
input_ids,
|
| 83 |
+
semantic_candidates,
|
| 84 |
+
seed,
|
| 85 |
+
device,
|
| 86 |
+
cfg_scale,
|
| 87 |
+
top_k,
|
| 88 |
+
reference_interval,
|
| 89 |
+
):
|
| 90 |
+
if semantic_candidates.ndim != 2 or semantic_candidates.shape[0] == 0:
|
| 91 |
+
raise ValueError("semantic candidates must have shape [frames, candidates]")
|
| 92 |
+
if semantic_candidates.shape[1] != REFERENCE_CANDIDATE_COUNT:
|
| 93 |
+
raise ValueError(f"semantic candidates must contain exactly {REFERENCE_CANDIDATE_COUNT} codes per frame")
|
| 94 |
+
if not 1 <= reference_interval <= 10:
|
| 95 |
+
raise ValueError("reference_interval must be between 1 and 10")
|
| 96 |
+
|
| 97 |
+
frame_count = semantic_candidates.shape[0]
|
| 98 |
prompt_tokens = int(input_ids.shape[1])
|
| 99 |
input_ids = input_ids.to(device)
|
| 100 |
+
semantic_candidates = semantic_candidates.to(device=device, dtype=torch.long)
|
| 101 |
execution_dtype = torch.bfloat16 if comfy.model_management.should_use_bf16(device) else torch.float32
|
| 102 |
unconditioned = input_ids.clone()
|
| 103 |
+
from comfy.ldm.minimax_music.prompt import SPECIAL_TOKEN_IDS
|
| 104 |
|
| 105 |
unconditioned[:, 1:-2] = SPECIAL_TOKEN_IDS["<|audio_cfg|>"]
|
| 106 |
text_ids = torch.cat((input_ids, unconditioned), dim=0)
|
| 107 |
+
|
| 108 |
+
def embed_text(token_ids):
|
| 109 |
+
if model.model.pruned_embedding:
|
| 110 |
+
return model.model.embed_tokens_prefill(token_ids, out_dtype=execution_dtype)
|
| 111 |
+
return model.model.embed_tokens(token_ids, out_dtype=execution_dtype)
|
| 112 |
+
|
| 113 |
+
past = model.model.init_kv_cache(2, prompt_tokens + frame_count + 2, device, execution_dtype)
|
| 114 |
+
output = model.model(None, embeds=embed_text(text_ids), past_key_values=past, dtype=execution_dtype)
|
| 115 |
last_hidden, past = output[0][:, -1], output[2]
|
| 116 |
from comfy.ldm.minimax_music.ar import derive_seed
|
| 117 |
|
|
|
|
| 145 |
output = model.model(None, embeds=feedback, past_key_values=past, dtype=execution_dtype)
|
| 146 |
last_hidden, past = output[0][:, -1], output[2]
|
| 147 |
|
| 148 |
+
all_semantic_candidates = torch.arange(16384, device=device, dtype=torch.long)
|
| 149 |
frames = []
|
| 150 |
+
for frame_index in range(frame_count):
|
| 151 |
comfy.model_management.throw_exception_if_processing_interrupted()
|
| 152 |
+
frame_candidates = (
|
| 153 |
+
semantic_candidates[frame_index]
|
| 154 |
+
if frame_index % reference_interval == 0
|
| 155 |
+
else all_semantic_candidates
|
| 156 |
+
)
|
| 157 |
+
c0 = _sample_semantic_candidates(
|
| 158 |
+
model,
|
| 159 |
+
last_hidden,
|
| 160 |
+
frame_candidates,
|
| 161 |
+
cfg_scale,
|
| 162 |
+
top_k,
|
| 163 |
+
generator,
|
| 164 |
+
).repeat(2)
|
| 165 |
+
c0_embed = model._embed_c0(c0, execution_dtype)
|
| 166 |
result: dict[str, torch.Tensor] = {}
|
| 167 |
|
| 168 |
+
def depth_core():
|
| 169 |
+
result["codes"], result["hidden"] = model._depth_codes(
|
| 170 |
+
last_hidden,
|
| 171 |
+
c0,
|
| 172 |
+
c0_embed,
|
| 173 |
+
generator,
|
| 174 |
+
execution_dtype,
|
| 175 |
+
cfg_scale,
|
| 176 |
+
top_k,
|
| 177 |
+
)
|
| 178 |
|
| 179 |
+
_run_depth(model, device, execution_dtype, depth_core)
|
| 180 |
frames.append(torch.cat((last_hidden[:1], result["hidden"]), dim=-1)[0].cpu())
|
| 181 |
+
if frame_index + 1 < frame_count:
|
| 182 |
+
feedback = model._embed_audio_frame(result["codes"], execution_dtype)
|
| 183 |
output = model.model(None, embeds=feedback, past_key_values=past, dtype=execution_dtype)
|
| 184 |
last_hidden, past = output[0][:, -1], output[2]
|
| 185 |
return torch.stack(frames)
|
|
|
|
| 189 |
|
| 190 |
|
| 191 |
def _encode_token_weights_with_reference(self, token_weight_pairs):
|
| 192 |
+
semantic_candidates = token_weight_pairs.get("minimax_reference_semantic_candidates")
|
| 193 |
+
if semantic_candidates is None:
|
| 194 |
return _ORIGINAL_ENCODE_TOKEN_WEIGHTS(self, token_weight_pairs)
|
| 195 |
token_ids = [token for token, _ in token_weight_pairs["minimax_music3"][0]]
|
| 196 |
input_ids = torch.tensor([token_ids], dtype=torch.long)
|
| 197 |
hidden = replay_reference_codes(
|
| 198 |
self,
|
| 199 |
input_ids,
|
| 200 |
+
semantic_candidates,
|
| 201 |
int(token_weight_pairs["seed"]),
|
| 202 |
self.execution_device,
|
| 203 |
float(token_weight_pairs["cfg_scale"]),
|
| 204 |
int(token_weight_pairs["top_k"]),
|
| 205 |
+
int(token_weight_pairs["minimax_reference_interval"]),
|
| 206 |
)
|
| 207 |
return hidden.unsqueeze(0), None, {}
|
| 208 |
|
|
|
|
| 215 |
class MiniMaxMusic3RVQReferenceEncoderLoader:
|
| 216 |
@classmethod
|
| 217 |
def INPUT_TYPES(cls):
|
| 218 |
+
encoders = [
|
| 219 |
+
name
|
| 220 |
+
for name in folder_paths.get_filename_list(ENCODER_FOLDER)
|
| 221 |
+
if name.lower().endswith(".safetensors")
|
| 222 |
+
]
|
| 223 |
dav_files = [name for name in folder_paths.get_filename_list("vae") if name.lower().endswith((".pth", ".pt"))]
|
| 224 |
return {"required": {"encoder": (encoders,), "dav": (dav_files,)}}
|
| 225 |
|
|
|
|
| 248 |
"caption": ("STRING", {"multiline": True, "dynamicPrompts": True}),
|
| 249 |
"lyrics": ("STRING", {"multiline": True, "dynamicPrompts": True}),
|
| 250 |
"seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
|
| 251 |
+
"reference_interval": ("INT", {"default": DEFAULT_REFERENCE_INTERVAL, "min": 1, "max": 10, "step": 1}),
|
| 252 |
},
|
| 253 |
"optional": {
|
| 254 |
"cfg_scale": ("FLOAT", {"default": CFG_SCALE, "min": 0.0, "max": 100.0, "step": 0.1}),
|
|
|
|
| 261 |
FUNCTION = "encode"
|
| 262 |
CATEGORY = "conditioning/minimax music"
|
| 263 |
|
| 264 |
+
def encode(
|
| 265 |
+
self,
|
| 266 |
+
clip,
|
| 267 |
+
reference_encoder,
|
| 268 |
+
audio,
|
| 269 |
+
caption,
|
| 270 |
+
lyrics,
|
| 271 |
+
seed,
|
| 272 |
+
reference_interval=DEFAULT_REFERENCE_INTERVAL,
|
| 273 |
+
cfg_scale=CFG_SCALE,
|
| 274 |
+
top_k=CFG_TOP_K,
|
| 275 |
+
):
|
| 276 |
+
codes, semantic_candidates = reference_encoder.predict_codes_with_semantic_candidates(
|
| 277 |
audio["waveform"],
|
| 278 |
int(audio["sample_rate"]),
|
| 279 |
+
semantic_top_k=REFERENCE_CANDIDATE_COUNT,
|
| 280 |
device=comfy.model_management.get_torch_device(),
|
| 281 |
)
|
| 282 |
tokens = clip.tokenize(
|
|
|
|
| 287 |
cfg_scale=cfg_scale,
|
| 288 |
top_k=top_k,
|
| 289 |
)
|
| 290 |
+
tokens["minimax_reference_semantic_candidates"] = semantic_candidates
|
| 291 |
+
tokens["minimax_reference_interval"] = reference_interval
|
| 292 |
conditioning = clip.encode_from_tokens_scheduled(tokens)
|
| 293 |
for cond in conditioning:
|
| 294 |
hidden = cond[0]
|
comfyui_workflow_example.json
CHANGED
|
@@ -369,7 +369,7 @@
|
|
| 369 |
],
|
| 370 |
"size": [
|
| 371 |
520,
|
| 372 |
-
|
| 373 |
],
|
| 374 |
"flags": {},
|
| 375 |
"order": 8,
|
|
@@ -418,6 +418,14 @@
|
|
| 418 |
"name": "seed"
|
| 419 |
}
|
| 420 |
},
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 421 |
{
|
| 422 |
"localized_name": "cfg_scale",
|
| 423 |
"name": "cfg_scale",
|
|
@@ -459,9 +467,10 @@
|
|
| 459 |
"Node name for S&R": "MiniMaxMusic3ReferenceAudioEncode"
|
| 460 |
},
|
| 461 |
"widgets_values": [
|
| 462 |
-
"
|
| 463 |
"[instrumental]",
|
| 464 |
299871780,
|
|
|
|
| 465 |
1.5,
|
| 466 |
50
|
| 467 |
]
|
|
|
|
| 369 |
],
|
| 370 |
"size": [
|
| 371 |
520,
|
| 372 |
+
560
|
| 373 |
],
|
| 374 |
"flags": {},
|
| 375 |
"order": 8,
|
|
|
|
| 418 |
"name": "seed"
|
| 419 |
}
|
| 420 |
},
|
| 421 |
+
{
|
| 422 |
+
"localized_name": "reference_interval",
|
| 423 |
+
"name": "reference_interval",
|
| 424 |
+
"type": "INT",
|
| 425 |
+
"widget": {
|
| 426 |
+
"name": "reference_interval"
|
| 427 |
+
}
|
| 428 |
+
},
|
| 429 |
{
|
| 430 |
"localized_name": "cfg_scale",
|
| 431 |
"name": "cfg_scale",
|
|
|
|
| 467 |
"Node name for S&R": "MiniMaxMusic3ReferenceAudioEncode"
|
| 468 |
},
|
| 469 |
"widgets_values": [
|
| 470 |
+
"1980s heavy metal cover, distorted electric guitars, driving bass, double-kick drums, dramatic arena production",
|
| 471 |
"[instrumental]",
|
| 472 |
299871780,
|
| 473 |
+
5,
|
| 474 |
1.5,
|
| 475 |
50
|
| 476 |
]
|
minimax_music3_reference_adapter.py
CHANGED
|
@@ -12,12 +12,12 @@
|
|
| 12 |
# See the License for the specific language governing permissions and
|
| 13 |
# limitations under the License.
|
| 14 |
|
| 15 |
-
"""Reference-audio
|
| 16 |
|
| 17 |
This file is self-contained. It loads a released SimpleTuner RVQ encoder,
|
| 18 |
encodes 44.1 kHz audio with the original DAV encoder, predicts eight RVQ codes
|
| 19 |
-
per 25 Hz frame, and
|
| 20 |
-
|
| 21 |
"""
|
| 22 |
|
| 23 |
from __future__ import annotations
|
|
@@ -52,7 +52,8 @@ AUDIO_CFG_TOKEN_ID = 151_654
|
|
| 52 |
SEMANTIC_VOCAB_SIZE = 16_384
|
| 53 |
AR_CFG_SCALE = 1.5
|
| 54 |
AR_TOP_K = 50
|
| 55 |
-
|
|
|
|
| 56 |
|
| 57 |
_SPECIAL_TAG_RE = re.compile(r"<\|([^|]*)\|>")
|
| 58 |
_LEADING_TAGS_RE = re.compile(r"^[ \t]*((?:\[[^\]]+\][ \t]*)+)")
|
|
@@ -406,26 +407,43 @@ def build_text_ids(tokenizer, prompt: str, lyrics: str, device: torch.device) ->
|
|
| 406 |
return torch.cat((input_ids, unconditional), dim=0).to(device)
|
| 407 |
|
| 408 |
|
| 409 |
-
def _sample_top_k(
|
|
|
|
|
|
|
|
|
|
|
|
|
| 410 |
values = torch.nan_to_num(logits.float(), nan=-1e9, posinf=1e9, neginf=-1e9)
|
| 411 |
-
threshold = torch.topk(values, min(
|
| 412 |
probabilities = torch.softmax(values.masked_fill(values < threshold, -torch.inf), dim=-1)
|
| 413 |
sample_device = generator.device if generator is not None else probabilities.device
|
| 414 |
return torch.multinomial(probabilities.to(sample_device), 1, generator=generator).squeeze(-1).to(values.device)
|
| 415 |
|
| 416 |
|
| 417 |
-
def
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 418 |
sequence = [depth_decoder.projection(hidden).unsqueeze(1)]
|
| 419 |
-
semantic = language_model.model.embed_tokens(
|
| 420 |
sequence.append(depth_decoder.projection(semantic).unsqueeze(1))
|
|
|
|
| 421 |
hidden_parts = []
|
| 422 |
-
for index in range(1,
|
| 423 |
depth_hidden = depth_decoder(torch.cat(sequence, dim=1))[:, -1]
|
| 424 |
-
hidden_parts.append(depth_hidden)
|
| 425 |
-
|
| 426 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 427 |
sequence.append(depth_decoder.projection(embedding).unsqueeze(1))
|
| 428 |
-
return torch.cat(hidden_parts, dim=-1)
|
| 429 |
|
| 430 |
|
| 431 |
def _embed_official_codes(language_model, depth_decoder, codes: torch.Tensor) -> torch.Tensor:
|
|
@@ -435,22 +453,53 @@ def _embed_official_codes(language_model, depth_decoder, codes: torch.Tensor) ->
|
|
| 435 |
return (semantic + acoustic.to(semantic.dtype)) * (codes.shape[-1] ** -0.5)
|
| 436 |
|
| 437 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 438 |
@torch.inference_mode()
|
| 439 |
def replay_codes_diffusers(
|
| 440 |
pipeline,
|
| 441 |
codes: torch.Tensor,
|
|
|
|
| 442 |
*,
|
| 443 |
prompt: str,
|
| 444 |
lyrics: str,
|
| 445 |
generator: torch.Generator | None = None,
|
|
|
|
|
|
|
|
|
|
| 446 |
) -> torch.Tensor:
|
| 447 |
-
"""
|
| 448 |
if codes.ndim != 2 or codes.shape[1] != 8 or codes.shape[0] == 0:
|
| 449 |
raise ValueError("codes must have shape [frames, 8]")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 450 |
language_model = pipeline.language_model
|
| 451 |
depth_decoder = pipeline.rvq_depth_decoder
|
| 452 |
device = next(language_model.parameters()).device
|
| 453 |
codes = codes.to(device=device, dtype=torch.long)
|
|
|
|
| 454 |
text_ids = build_text_ids(pipeline.tokenizer, prompt, lyrics, device)
|
| 455 |
text_output = language_model.model(inputs_embeds=language_model.model.embed_tokens(text_ids), use_cache=True)
|
| 456 |
past = text_output.past_key_values
|
|
@@ -461,56 +510,64 @@ def replay_codes_diffusers(
|
|
| 461 |
vocab_mask[AUDIO_END_TOKEN_ID] = False
|
| 462 |
logits = language_model.lm_head(last_hidden).float().masked_fill(vocab_mask, -torch.inf)
|
| 463 |
conditioned, unconditioned = logits[:1], logits[1:2]
|
| 464 |
-
guided = unconditioned + (conditioned - unconditioned) *
|
| 465 |
-
threshold = torch.topk(conditioned,
|
| 466 |
-
warmup_token = _sample_top_k(guided.masked_fill(conditioned < threshold, -torch.inf), generator)
|
| 467 |
if int(warmup_token.item()) == AUDIO_END_TOKEN_ID:
|
| 468 |
raise ValueError("the selected seed ended during the required AR warm-up frame")
|
| 469 |
warmup_semantic = (warmup_token - AUDIO_CODE_OFFSET).repeat(2)
|
| 470 |
|
| 471 |
-
|
| 472 |
-
|
| 473 |
-
|
| 474 |
-
|
| 475 |
-
|
| 476 |
-
|
| 477 |
-
|
| 478 |
-
|
| 479 |
-
|
| 480 |
-
warmup_codes.append(code)
|
| 481 |
-
if index < 7:
|
| 482 |
-
embedding = depth_decoder.audio_embeddings(code + (index - 1) * 1024)
|
| 483 |
-
sequence.append(depth_decoder.projection(embedding).unsqueeze(1))
|
| 484 |
-
warmup_codes = torch.stack(warmup_codes, dim=1)
|
| 485 |
output = language_model.model(
|
| 486 |
inputs_embeds=_embed_official_codes(language_model, depth_decoder, warmup_codes).unsqueeze(1),
|
| 487 |
past_key_values=past,
|
| 488 |
use_cache=True,
|
| 489 |
)
|
| 490 |
past = output.past_key_values
|
| 491 |
-
|
| 492 |
-
|
| 493 |
-
|
| 494 |
-
|
| 495 |
-
|
| 496 |
-
|
| 497 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 498 |
output = language_model.model(
|
| 499 |
-
inputs_embeds=
|
| 500 |
past_key_values=past,
|
| 501 |
use_cache=True,
|
| 502 |
)
|
| 503 |
past = output.past_key_values
|
| 504 |
-
|
| 505 |
-
|
| 506 |
-
repeated_codes = codes.unsqueeze(0).repeat(2, 1, 1).reshape(-1, 8)
|
| 507 |
-
depth_hidden = _official_depth_hidden(
|
| 508 |
-
language_model,
|
| 509 |
-
depth_decoder,
|
| 510 |
-
global_hidden.reshape(-1, global_hidden.shape[-1]),
|
| 511 |
-
repeated_codes,
|
| 512 |
-
).view(2, codes.shape[0], -1)
|
| 513 |
-
return torch.cat((global_hidden[:1], depth_hidden[:1]), dim=-1).cpu()
|
| 514 |
|
| 515 |
|
| 516 |
class MiniMaxMusic3ReferenceAdapter:
|
|
@@ -577,14 +634,19 @@ class MiniMaxMusic3ReferenceAdapter:
|
|
| 577 |
return waveform.float()
|
| 578 |
|
| 579 |
@torch.inference_mode()
|
| 580 |
-
def
|
| 581 |
self,
|
| 582 |
waveform: torch.Tensor,
|
| 583 |
sample_rate: int,
|
| 584 |
*,
|
| 585 |
device: str | torch.device | None = None,
|
| 586 |
encoder_dtype: torch.dtype | None = None,
|
| 587 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 588 |
waveform = self._resample(waveform, sample_rate)
|
| 589 |
original_samples = waveform.shape[-1]
|
| 590 |
if original_samples < SAMPLE_RATE // FRAME_RATE:
|
|
@@ -616,6 +678,9 @@ class MiniMaxMusic3ReferenceAdapter:
|
|
| 616 |
regular_starts = [0]
|
| 617 |
|
| 618 |
predictions = torch.empty((frame_count, 8), dtype=torch.long)
|
|
|
|
|
|
|
|
|
|
| 619 |
assigned = torch.zeros(frame_count, dtype=torch.bool)
|
| 620 |
self.rvq_encoder.to(device=device, dtype=encoder_dtype)
|
| 621 |
autocast = torch.autocast(device.type, dtype=encoder_dtype) if device.type == "cuda" else nullcontext()
|
|
@@ -628,14 +693,56 @@ class MiniMaxMusic3ReferenceAdapter:
|
|
| 628 |
pool = build_pool_matrix(local_bounds).to(device)
|
| 629 |
logits = self.rvq_encoder(window_latents.unsqueeze(0), pool.unsqueeze(0))
|
| 630 |
predicted = torch.stack([head.argmax(dim=-1)[0] for head in logits], dim=-1).cpu()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 631 |
take = ~assigned[frame_start:frame_end]
|
| 632 |
predictions[frame_start:frame_end][take] = predicted[take]
|
|
|
|
|
|
|
| 633 |
assigned[frame_start:frame_end][take] = True
|
| 634 |
self.rvq_encoder.to("cpu")
|
| 635 |
if not assigned.all():
|
| 636 |
raise RuntimeError("RVQ window inference did not cover every reference frame")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 637 |
return predictions
|
| 638 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 639 |
def encode_reference(
|
| 640 |
self,
|
| 641 |
pipeline,
|
|
@@ -646,14 +753,26 @@ class MiniMaxMusic3ReferenceAdapter:
|
|
| 646 |
lyrics: str,
|
| 647 |
generator: torch.Generator | None = None,
|
| 648 |
device: str | torch.device | None = None,
|
|
|
|
|
|
|
|
|
|
| 649 |
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 650 |
-
codes = self.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 651 |
frame_hiddens = replay_codes_diffusers(
|
| 652 |
pipeline,
|
| 653 |
codes,
|
|
|
|
| 654 |
prompt=prompt,
|
| 655 |
lyrics=lyrics,
|
| 656 |
generator=generator,
|
|
|
|
|
|
|
|
|
|
| 657 |
)
|
| 658 |
return frame_hiddens, codes
|
| 659 |
|
|
|
|
| 12 |
# See the License for the specific language governing permissions and
|
| 13 |
# limitations under the License.
|
| 14 |
|
| 15 |
+
"""Reference-audio constrained generation for MiniMax Music 3.
|
| 16 |
|
| 17 |
This file is self-contained. It loads a released SimpleTuner RVQ encoder,
|
| 18 |
encodes 44.1 kHz audio with the original DAV encoder, predicts eight RVQ codes
|
| 19 |
+
per 25 Hz frame, and periodically constrains the official MiniMax Music 3
|
| 20 |
+
language model to the encoder's top-5 semantic candidates.
|
| 21 |
"""
|
| 22 |
|
| 23 |
from __future__ import annotations
|
|
|
|
| 52 |
SEMANTIC_VOCAB_SIZE = 16_384
|
| 53 |
AR_CFG_SCALE = 1.5
|
| 54 |
AR_TOP_K = 50
|
| 55 |
+
REFERENCE_CANDIDATE_COUNT = 5
|
| 56 |
+
DEFAULT_REFERENCE_INTERVAL = 5
|
| 57 |
|
| 58 |
_SPECIAL_TAG_RE = re.compile(r"<\|([^|]*)\|>")
|
| 59 |
_LEADING_TAGS_RE = re.compile(r"^[ \t]*((?:\[[^\]]+\][ \t]*)+)")
|
|
|
|
| 407 |
return torch.cat((input_ids, unconditional), dim=0).to(device)
|
| 408 |
|
| 409 |
|
| 410 |
+
def _sample_top_k(
|
| 411 |
+
logits: torch.Tensor,
|
| 412 |
+
generator: torch.Generator | None,
|
| 413 |
+
top_k: int = AR_TOP_K,
|
| 414 |
+
) -> torch.Tensor:
|
| 415 |
values = torch.nan_to_num(logits.float(), nan=-1e9, posinf=1e9, neginf=-1e9)
|
| 416 |
+
threshold = torch.topk(values, min(top_k, values.shape[-1]), dim=-1).values[..., -1, None]
|
| 417 |
probabilities = torch.softmax(values.masked_fill(values < threshold, -torch.inf), dim=-1)
|
| 418 |
sample_device = generator.device if generator is not None else probabilities.device
|
| 419 |
return torch.multinomial(probabilities.to(sample_device), 1, generator=generator).squeeze(-1).to(values.device)
|
| 420 |
|
| 421 |
|
| 422 |
+
def _sample_official_depth_codes(
|
| 423 |
+
language_model,
|
| 424 |
+
depth_decoder,
|
| 425 |
+
hidden: torch.Tensor,
|
| 426 |
+
semantic_codes: torch.Tensor,
|
| 427 |
+
generator: torch.Generator | None,
|
| 428 |
+
cfg_scale: float = AR_CFG_SCALE,
|
| 429 |
+
top_k: int = AR_TOP_K,
|
| 430 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 431 |
sequence = [depth_decoder.projection(hidden).unsqueeze(1)]
|
| 432 |
+
semantic = language_model.model.embed_tokens(semantic_codes + AUDIO_CODE_OFFSET)
|
| 433 |
sequence.append(depth_decoder.projection(semantic).unsqueeze(1))
|
| 434 |
+
codes = [semantic_codes]
|
| 435 |
hidden_parts = []
|
| 436 |
+
for index in range(1, 8):
|
| 437 |
depth_hidden = depth_decoder(torch.cat(sequence, dim=1))[:, -1]
|
| 438 |
+
hidden_parts.append(depth_hidden[:1])
|
| 439 |
+
logits = depth_decoder.audio_heads[index - 1](depth_hidden).float()
|
| 440 |
+
guided = logits[1:2] + (logits[:1] - logits[1:2]) * cfg_scale
|
| 441 |
+
code = _sample_top_k(guided, generator, top_k).repeat(2)
|
| 442 |
+
codes.append(code)
|
| 443 |
+
if index < 7:
|
| 444 |
+
embedding = depth_decoder.audio_embeddings(code + (index - 1) * 1024)
|
| 445 |
sequence.append(depth_decoder.projection(embedding).unsqueeze(1))
|
| 446 |
+
return torch.stack(codes, dim=1), torch.cat(hidden_parts, dim=-1)
|
| 447 |
|
| 448 |
|
| 449 |
def _embed_official_codes(language_model, depth_decoder, codes: torch.Tensor) -> torch.Tensor:
|
|
|
|
| 453 |
return (semantic + acoustic.to(semantic.dtype)) * (codes.shape[-1] ** -0.5)
|
| 454 |
|
| 455 |
|
| 456 |
+
def _sample_semantic_candidates_diffusers(
|
| 457 |
+
language_model,
|
| 458 |
+
hidden: torch.Tensor,
|
| 459 |
+
semantic_candidates: torch.Tensor,
|
| 460 |
+
cfg_scale: float,
|
| 461 |
+
top_k: int,
|
| 462 |
+
generator: torch.Generator | None,
|
| 463 |
+
) -> torch.Tensor:
|
| 464 |
+
token_candidates = semantic_candidates.to(device=hidden.device, dtype=torch.long) + AUDIO_CODE_OFFSET
|
| 465 |
+
token_candidates = token_candidates.unsqueeze(0)
|
| 466 |
+
logits = language_model.lm_head(hidden).float()
|
| 467 |
+
conditioned = logits[:1].gather(-1, token_candidates)
|
| 468 |
+
unconditioned = logits[1:2].gather(-1, token_candidates)
|
| 469 |
+
guided = unconditioned + (conditioned - unconditioned) * cfg_scale
|
| 470 |
+
selected_index = _sample_top_k(guided, generator, min(top_k, guided.shape[-1])).view(1, 1)
|
| 471 |
+
return token_candidates.gather(-1, selected_index).squeeze(-1) - AUDIO_CODE_OFFSET
|
| 472 |
+
|
| 473 |
+
|
| 474 |
@torch.inference_mode()
|
| 475 |
def replay_codes_diffusers(
|
| 476 |
pipeline,
|
| 477 |
codes: torch.Tensor,
|
| 478 |
+
semantic_candidates: torch.Tensor,
|
| 479 |
*,
|
| 480 |
prompt: str,
|
| 481 |
lyrics: str,
|
| 482 |
generator: torch.Generator | None = None,
|
| 483 |
+
reference_interval: int = DEFAULT_REFERENCE_INTERVAL,
|
| 484 |
+
cfg_scale: float = AR_CFG_SCALE,
|
| 485 |
+
top_k: int = AR_TOP_K,
|
| 486 |
) -> torch.Tensor:
|
| 487 |
+
"""Generate a fixed-length rollout with periodic top-5 semantic constraints."""
|
| 488 |
if codes.ndim != 2 or codes.shape[1] != 8 or codes.shape[0] == 0:
|
| 489 |
raise ValueError("codes must have shape [frames, 8]")
|
| 490 |
+
if semantic_candidates.ndim != 2 or semantic_candidates.shape != (codes.shape[0], REFERENCE_CANDIDATE_COUNT):
|
| 491 |
+
raise ValueError(
|
| 492 |
+
f"semantic_candidates must have shape [frames, {REFERENCE_CANDIDATE_COUNT}]"
|
| 493 |
+
)
|
| 494 |
+
if not 1 <= reference_interval <= 10:
|
| 495 |
+
raise ValueError("reference_interval must be between 1 and 10")
|
| 496 |
+
if not 1 <= top_k <= SEMANTIC_VOCAB_SIZE:
|
| 497 |
+
raise ValueError(f"top_k must be between 1 and {SEMANTIC_VOCAB_SIZE}")
|
| 498 |
language_model = pipeline.language_model
|
| 499 |
depth_decoder = pipeline.rvq_depth_decoder
|
| 500 |
device = next(language_model.parameters()).device
|
| 501 |
codes = codes.to(device=device, dtype=torch.long)
|
| 502 |
+
semantic_candidates = semantic_candidates.to(device=device, dtype=torch.long)
|
| 503 |
text_ids = build_text_ids(pipeline.tokenizer, prompt, lyrics, device)
|
| 504 |
text_output = language_model.model(inputs_embeds=language_model.model.embed_tokens(text_ids), use_cache=True)
|
| 505 |
past = text_output.past_key_values
|
|
|
|
| 510 |
vocab_mask[AUDIO_END_TOKEN_ID] = False
|
| 511 |
logits = language_model.lm_head(last_hidden).float().masked_fill(vocab_mask, -torch.inf)
|
| 512 |
conditioned, unconditioned = logits[:1], logits[1:2]
|
| 513 |
+
guided = unconditioned + (conditioned - unconditioned) * cfg_scale
|
| 514 |
+
threshold = torch.topk(conditioned, top_k, dim=-1).values[..., -1, None]
|
| 515 |
+
warmup_token = _sample_top_k(guided.masked_fill(conditioned < threshold, -torch.inf), generator, top_k)
|
| 516 |
if int(warmup_token.item()) == AUDIO_END_TOKEN_ID:
|
| 517 |
raise ValueError("the selected seed ended during the required AR warm-up frame")
|
| 518 |
warmup_semantic = (warmup_token - AUDIO_CODE_OFFSET).repeat(2)
|
| 519 |
|
| 520 |
+
warmup_codes, _ = _sample_official_depth_codes(
|
| 521 |
+
language_model,
|
| 522 |
+
depth_decoder,
|
| 523 |
+
last_hidden,
|
| 524 |
+
warmup_semantic,
|
| 525 |
+
generator,
|
| 526 |
+
cfg_scale=cfg_scale,
|
| 527 |
+
top_k=top_k,
|
| 528 |
+
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 529 |
output = language_model.model(
|
| 530 |
inputs_embeds=_embed_official_codes(language_model, depth_decoder, warmup_codes).unsqueeze(1),
|
| 531 |
past_key_values=past,
|
| 532 |
use_cache=True,
|
| 533 |
)
|
| 534 |
past = output.past_key_values
|
| 535 |
+
last_hidden = output.last_hidden_state[:, -1]
|
| 536 |
+
all_semantic_candidates = torch.arange(SEMANTIC_VOCAB_SIZE, device=device, dtype=torch.long)
|
| 537 |
+
hidden_frames = []
|
| 538 |
+
for frame_index in range(codes.shape[0]):
|
| 539 |
+
frame_candidates = (
|
| 540 |
+
semantic_candidates[frame_index]
|
| 541 |
+
if frame_index % reference_interval == 0
|
| 542 |
+
else all_semantic_candidates
|
| 543 |
+
)
|
| 544 |
+
semantic_code = _sample_semantic_candidates_diffusers(
|
| 545 |
+
language_model,
|
| 546 |
+
last_hidden,
|
| 547 |
+
frame_candidates,
|
| 548 |
+
cfg_scale,
|
| 549 |
+
top_k,
|
| 550 |
+
generator,
|
| 551 |
+
).repeat(2)
|
| 552 |
+
sampled_codes, depth_hidden = _sample_official_depth_codes(
|
| 553 |
+
language_model,
|
| 554 |
+
depth_decoder,
|
| 555 |
+
last_hidden,
|
| 556 |
+
semantic_code,
|
| 557 |
+
generator,
|
| 558 |
+
cfg_scale=cfg_scale,
|
| 559 |
+
top_k=top_k,
|
| 560 |
+
)
|
| 561 |
+
hidden_frames.append(torch.cat((last_hidden[:1], depth_hidden), dim=-1).cpu())
|
| 562 |
+
if frame_index + 1 < codes.shape[0]:
|
| 563 |
output = language_model.model(
|
| 564 |
+
inputs_embeds=_embed_official_codes(language_model, depth_decoder, sampled_codes).unsqueeze(1),
|
| 565 |
past_key_values=past,
|
| 566 |
use_cache=True,
|
| 567 |
)
|
| 568 |
past = output.past_key_values
|
| 569 |
+
last_hidden = output.last_hidden_state[:, -1]
|
| 570 |
+
return torch.cat(hidden_frames, dim=0).unsqueeze(0)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 571 |
|
| 572 |
|
| 573 |
class MiniMaxMusic3ReferenceAdapter:
|
|
|
|
| 634 |
return waveform.float()
|
| 635 |
|
| 636 |
@torch.inference_mode()
|
| 637 |
+
def _predict_codes(
|
| 638 |
self,
|
| 639 |
waveform: torch.Tensor,
|
| 640 |
sample_rate: int,
|
| 641 |
*,
|
| 642 |
device: str | torch.device | None = None,
|
| 643 |
encoder_dtype: torch.dtype | None = None,
|
| 644 |
+
semantic_top_k: int | None = None,
|
| 645 |
+
) -> tuple[torch.Tensor, torch.Tensor | None]:
|
| 646 |
+
if semantic_top_k is not None and not 1 <= semantic_top_k <= self.rvq_encoder.config.codebook_vocab_sizes[0]:
|
| 647 |
+
raise ValueError(
|
| 648 |
+
f"semantic_top_k must be between 1 and {self.rvq_encoder.config.codebook_vocab_sizes[0]}"
|
| 649 |
+
)
|
| 650 |
waveform = self._resample(waveform, sample_rate)
|
| 651 |
original_samples = waveform.shape[-1]
|
| 652 |
if original_samples < SAMPLE_RATE // FRAME_RATE:
|
|
|
|
| 678 |
regular_starts = [0]
|
| 679 |
|
| 680 |
predictions = torch.empty((frame_count, 8), dtype=torch.long)
|
| 681 |
+
semantic_candidates = (
|
| 682 |
+
torch.empty((frame_count, semantic_top_k), dtype=torch.long) if semantic_top_k is not None else None
|
| 683 |
+
)
|
| 684 |
assigned = torch.zeros(frame_count, dtype=torch.bool)
|
| 685 |
self.rvq_encoder.to(device=device, dtype=encoder_dtype)
|
| 686 |
autocast = torch.autocast(device.type, dtype=encoder_dtype) if device.type == "cuda" else nullcontext()
|
|
|
|
| 693 |
pool = build_pool_matrix(local_bounds).to(device)
|
| 694 |
logits = self.rvq_encoder(window_latents.unsqueeze(0), pool.unsqueeze(0))
|
| 695 |
predicted = torch.stack([head.argmax(dim=-1)[0] for head in logits], dim=-1).cpu()
|
| 696 |
+
window_semantic_candidates = (
|
| 697 |
+
logits[0].topk(semantic_top_k, dim=-1).indices[0].cpu()
|
| 698 |
+
if semantic_top_k is not None
|
| 699 |
+
else None
|
| 700 |
+
)
|
| 701 |
take = ~assigned[frame_start:frame_end]
|
| 702 |
predictions[frame_start:frame_end][take] = predicted[take]
|
| 703 |
+
if semantic_candidates is not None:
|
| 704 |
+
semantic_candidates[frame_start:frame_end][take] = window_semantic_candidates[take]
|
| 705 |
assigned[frame_start:frame_end][take] = True
|
| 706 |
self.rvq_encoder.to("cpu")
|
| 707 |
if not assigned.all():
|
| 708 |
raise RuntimeError("RVQ window inference did not cover every reference frame")
|
| 709 |
+
return predictions, semantic_candidates
|
| 710 |
+
|
| 711 |
+
def predict_codes(
|
| 712 |
+
self,
|
| 713 |
+
waveform: torch.Tensor,
|
| 714 |
+
sample_rate: int,
|
| 715 |
+
*,
|
| 716 |
+
device: str | torch.device | None = None,
|
| 717 |
+
encoder_dtype: torch.dtype | None = None,
|
| 718 |
+
) -> torch.Tensor:
|
| 719 |
+
predictions, _ = self._predict_codes(
|
| 720 |
+
waveform,
|
| 721 |
+
sample_rate,
|
| 722 |
+
device=device,
|
| 723 |
+
encoder_dtype=encoder_dtype,
|
| 724 |
+
)
|
| 725 |
return predictions
|
| 726 |
|
| 727 |
+
def predict_codes_with_semantic_candidates(
|
| 728 |
+
self,
|
| 729 |
+
waveform: torch.Tensor,
|
| 730 |
+
sample_rate: int,
|
| 731 |
+
*,
|
| 732 |
+
semantic_top_k: int = 5,
|
| 733 |
+
device: str | torch.device | None = None,
|
| 734 |
+
encoder_dtype: torch.dtype | None = None,
|
| 735 |
+
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 736 |
+
predictions, semantic_candidates = self._predict_codes(
|
| 737 |
+
waveform,
|
| 738 |
+
sample_rate,
|
| 739 |
+
device=device,
|
| 740 |
+
encoder_dtype=encoder_dtype,
|
| 741 |
+
semantic_top_k=semantic_top_k,
|
| 742 |
+
)
|
| 743 |
+
assert semantic_candidates is not None
|
| 744 |
+
return predictions, semantic_candidates
|
| 745 |
+
|
| 746 |
def encode_reference(
|
| 747 |
self,
|
| 748 |
pipeline,
|
|
|
|
| 753 |
lyrics: str,
|
| 754 |
generator: torch.Generator | None = None,
|
| 755 |
device: str | torch.device | None = None,
|
| 756 |
+
reference_interval: int = DEFAULT_REFERENCE_INTERVAL,
|
| 757 |
+
cfg_scale: float = AR_CFG_SCALE,
|
| 758 |
+
top_k: int = AR_TOP_K,
|
| 759 |
) -> tuple[torch.Tensor, torch.Tensor]:
|
| 760 |
+
codes, semantic_candidates = self.predict_codes_with_semantic_candidates(
|
| 761 |
+
waveform,
|
| 762 |
+
sample_rate,
|
| 763 |
+
semantic_top_k=REFERENCE_CANDIDATE_COUNT,
|
| 764 |
+
device=device,
|
| 765 |
+
)
|
| 766 |
frame_hiddens = replay_codes_diffusers(
|
| 767 |
pipeline,
|
| 768 |
codes,
|
| 769 |
+
semantic_candidates,
|
| 770 |
prompt=prompt,
|
| 771 |
lyrics=lyrics,
|
| 772 |
generator=generator,
|
| 773 |
+
reference_interval=reference_interval,
|
| 774 |
+
cfg_scale=cfg_scale,
|
| 775 |
+
top_k=top_k,
|
| 776 |
)
|
| 777 |
return frame_hiddens, codes
|
| 778 |
|