Minimax Music 3: optional NVFP4 quantization (MM3_NVFP4=1)

Adds an opt-in load-time conversion of the int8-quantized MiniMax Music 3 checkpoints (AR/TE and DiT) to NVFP4 using comfy-kitchen's TensorCoreNVFP4Layout. NVFP4 keeps the same 2D quantized-weight contract with half the storage bytes, so the AR model's resident set fits VRAM on 8GB cards and the aimdo stops paging the weights mid-generation.

The conversion dequantizes int8 layers (tensorwise, with or without convrot) and re-quantizes them in place; non-quantized layers are untouched. Gated behind the MM3_NVFP4 environment variable; about 1-2 minutes of extra load time, which can be avoided by pre-baking the converted checkpoints.
This commit is contained in:
F1LIP3 2026-08-15 14:22:32 -03:00
parent 0f1fa67ad8
commit 5ae70eec8a
2 changed files with 38 additions and 0 deletions

View File

@ -2267,6 +2267,9 @@ def load_diffusion_model_state_dict(sd, model_options={}, metadata=None, disable
parameters = comfy.utils.calculate_parameters(sd)
weight_dtype = comfy.utils.weight_dtype(sd)
if os.environ.get("MM3_NVFP4") and metadata is not None and metadata.get("comfy_model") == "minimax_music3_dit":
comfy.text_encoders.minimax_music._to_nvfp4(sd)
load_device = model_options.get("load_device", model_management.get_torch_device())
model_config = model_detection.model_config_from_unet(sd, "", metadata=metadata)

View File

@ -1,3 +1,6 @@
import json
import os
import torch
from tokenizers import Tokenizer
@ -6,6 +9,36 @@ from comfy.ldm.minimax_music.ar import CFG_SCALE, CFG_TOP_K, MAX_AUDIO_FRAMES, M
from comfy.ldm.minimax_music.prompt import SPECIAL_TOKEN_IDS, build_prompt
def _to_nvfp4(sd):
"""Rewrite int8+convrot quantized layers to NVFP4 in place (MM3_NVFP4=1).
NVFP4 keeps the same 2D quantized-weight contract with ~half the storage
bytes, so the AR model's resident set fits VRAM and aimdo stops paging.
"""
import comfy_kitchen as ck
from comfy_kitchen.tensor.nvfp4 import TensorCoreNVFP4Layout
dtype_code = ck.DTYPE_TO_CODE[torch.float32]
for key in list(sd):
if not key.endswith(".comfy_quant"):
continue
conf = json.loads(sd[key].numpy().tobytes())
if conf.get("format") != "int8_tensorwise":
continue
prefix = key[:-len(".comfy_quant")]
params_conf = conf.get("params", {})
groupsize = int(conf.get("convrot_groupsize", params_conf.get("convrot_groupsize", 256)))
if conf.get("convrot", params_conf.get("convrot", False)):
fp = torch.ops.comfy_kitchen.dequantize_int8_convrot_weight_dtype(sd[prefix + ".weight"], sd[prefix + ".weight_scale"], groupsize, dtype_code)
else:
fp = ck.dequantize_int8_simple_dtype(sd[prefix + ".weight"], sd[prefix + ".weight_scale"], dtype_code)
qdata, params = TensorCoreNVFP4Layout.quantize(fp)
sd[prefix + ".weight"] = qdata
sd[prefix + ".weight_scale"] = params.block_scale
sd[prefix + ".weight_scale_2"] = params.scale
sd[key] = torch.tensor(list(b'{"format": "nvfp4"}'), dtype=torch.uint8)
MODEL_CONFIG = {
"vocab_size": 200000,
"hidden_size": 4096,
@ -99,6 +132,8 @@ class MiniMaxMusic3TEModel(MiniMaxMusic3AR):
return hidden.unsqueeze(0), None, {}
def load_state_dict(self, state_dict, strict=True, assign=False):
if os.environ.get("MM3_NVFP4"):
_to_nvfp4(state_dict)
if self.model.pruned_embedding is None:
self.model.pruned_embedding = "model.embed_tokens_prefill.weight" in state_dict
if self.model.pruned_embedding: