PyTorch
Diffusers
audio
music
minimax-music-3
rvq
reference-audio
comfyui
bghira commited on
Commit
0e3682a
·
verified ·
1 Parent(s): cb1b96a

Keep only periodic c0 reference constraints

Browse files
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
- | Version | Parameters | Experiment | Replay cosine |
33
  |---|---:|---|---:|
34
- | v1 | 40,978,944 | Baseline; eight independent heads | 0.7624* |
35
- | v2 | 154,736,064 | Wider shared encoder | 0.7698 |
36
- | v3 | 154,736,064 | v2 plus training-only MERT alignment | 0.7703 |
37
- | **v4** | 169,008,576 | Causal acoustic decoder across codebook depth | **0.8748** |
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. Load [`comfyui_workflow_example.json`](comfyui_workflow_example.json). Upload a reference audio file. Select v4 in **MiniMax Music3 RVQ Reference Encoder Loader**.
 
 
 
 
 
 
 
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 tested graph used ComfyUI 0.33.0, one NVIDIA L40S, the pruned int8 text encoder, the fp16 diffusion model, a 32.28-second held-out reference, five Euler steps, and the v4 encoder. It completed and produced a full-length non-silent stereo FLAC. Use 30 steps for normal output.
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 replay still needs the official MiniMax language model and RVQ depth decoder.
213
- - v4 uses greedy code selection. Other decoding strategies remain untested.
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
- folder_paths.add_model_folder_path(ENCODER_FOLDER, str(Path(folder_paths.models_dir) / ENCODER_FOLDER), is_default=True)
 
 
 
 
 
 
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(model: MiniMaxMusic3AR, input_ids, codes, seed, device, cfg_scale, top_k):
76
- if codes.ndim != 2 or codes.shape[1] != model.num_codebooks or codes.shape[0] == 0:
77
- raise ValueError(f"reference codes must have shape [frames, {model.num_codebooks}]")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
78
  prompt_tokens = int(input_ids.shape[1])
79
  input_ids = input_ids.to(device)
80
- codes = codes.to(device=device, dtype=torch.long)
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 AUDIO_CODE_OFFSET, SPECIAL_TOKEN_IDS
84
 
85
  unconditioned[:, 1:-2] = SPECIAL_TOKEN_IDS["<|audio_cfg|>"]
86
  text_ids = torch.cat((input_ids, unconditioned), dim=0)
87
- if model.model.pruned_embedding:
88
- text_embeds = model.model.embed_tokens_prefill(text_ids, out_dtype=execution_dtype)
89
- else:
90
- text_embeds = model.model.embed_tokens(text_ids, out_dtype=execution_dtype)
91
- past = model.model.init_kv_cache(2, prompt_tokens + codes.shape[0] + 2, device, execution_dtype)
92
- output = model.model(None, embeds=text_embeds, past_key_values=past, dtype=execution_dtype)
 
 
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(codes.shape[0]):
128
  comfy.model_management.throw_exception_if_processing_interrupted()
129
- frame_codes = codes[frame_index].unsqueeze(0).repeat(2, 1)
 
 
 
 
 
 
 
 
 
 
 
 
 
130
  result: dict[str, torch.Tensor] = {}
131
 
132
- def teacher_core():
133
- result["hidden"] = _teacher_depth_hidden(model, last_hidden, frame_codes, execution_dtype)
 
 
 
 
 
 
 
 
134
 
135
- _run_depth(model, device, execution_dtype, teacher_core)
136
  frames.append(torch.cat((last_hidden[:1], result["hidden"]), dim=-1)[0].cpu())
137
- if frame_index + 1 < codes.shape[0]:
138
- feedback = model._embed_audio_frame(frame_codes, execution_dtype)
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
- codes = token_weight_pairs.get("minimax_reference_codes")
149
- if codes is None:
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
- codes,
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 = [name for name in folder_paths.get_filename_list(ENCODER_FOLDER) if name.lower().endswith(".safetensors")]
 
 
 
 
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(self, clip, reference_encoder, audio, caption, lyrics, seed, cfg_scale=CFG_SCALE, top_k=CFG_TOP_K):
215
- codes = reference_encoder.predict_codes(
 
 
 
 
 
 
 
 
 
 
 
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["minimax_reference_codes"] = codes
 
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
- 640
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
- "rock",
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 conditioning 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 teacher-forces those codes through the official MiniMax
20
- Music 3 language-model path.
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
- LM_BLOCK_FRAMES = 256
 
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(logits: torch.Tensor, generator: torch.Generator | None) -> torch.Tensor:
 
 
 
 
410
  values = torch.nan_to_num(logits.float(), nan=-1e9, posinf=1e9, neginf=-1e9)
411
- threshold = torch.topk(values, min(AR_TOP_K, values.shape[-1]), dim=-1).values[..., -1, None]
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 _official_depth_hidden(language_model, depth_decoder, hidden: torch.Tensor, codes: torch.Tensor) -> torch.Tensor:
 
 
 
 
 
 
 
 
418
  sequence = [depth_decoder.projection(hidden).unsqueeze(1)]
419
- semantic = language_model.model.embed_tokens(codes[:, 0] + AUDIO_CODE_OFFSET)
420
  sequence.append(depth_decoder.projection(semantic).unsqueeze(1))
 
421
  hidden_parts = []
422
- for index in range(1, codes.shape[1]):
423
  depth_hidden = depth_decoder(torch.cat(sequence, dim=1))[:, -1]
424
- hidden_parts.append(depth_hidden)
425
- if index < codes.shape[1] - 1:
426
- embedding = depth_decoder.audio_embeddings(codes[:, index] + (index - 1) * 1024)
 
 
 
 
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
- """Replay predicted codes through official Diffusers LM components."""
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) * AR_CFG_SCALE
465
- threshold = torch.topk(conditioned, AR_TOP_K, dim=-1).values[..., -1, None]
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
- sequence = [depth_decoder.projection(last_hidden).unsqueeze(1)]
472
- semantic_embed = language_model.model.embed_tokens(warmup_semantic + AUDIO_CODE_OFFSET)
473
- sequence.append(depth_decoder.projection(semantic_embed).unsqueeze(1))
474
- warmup_codes = [warmup_semantic]
475
- for index in range(1, 8):
476
- depth_hidden = depth_decoder(torch.cat(sequence, dim=1))[:, -1]
477
- depth_logits = depth_decoder.audio_heads[index - 1](depth_hidden).float()
478
- depth_guided = depth_logits[1:2] + (depth_logits[:1] - depth_logits[1:2]) * AR_CFG_SCALE
479
- code = _sample_top_k(depth_guided, generator).repeat(2)
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
- first_hidden = output.last_hidden_state[:, -1]
492
-
493
- hidden_parts = [first_hidden.unsqueeze(1)]
494
- if codes.shape[0] > 1:
495
- feedback_codes = codes[:-1].unsqueeze(0).repeat(2, 1, 1)
496
- feedback = _embed_official_codes(language_model, depth_decoder, feedback_codes)
497
- for start in range(0, feedback.shape[1], LM_BLOCK_FRAMES):
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
498
  output = language_model.model(
499
- inputs_embeds=feedback[:, start : start + LM_BLOCK_FRAMES],
500
  past_key_values=past,
501
  use_cache=True,
502
  )
503
  past = output.past_key_values
504
- hidden_parts.append(output.last_hidden_state)
505
- global_hidden = torch.cat(hidden_parts, dim=1)
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 predict_codes(
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
- ) -> torch.Tensor:
 
 
 
 
 
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.predict_codes(waveform, sample_rate, device=device)
 
 
 
 
 
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