Add support for MiniMax H3 T2V and I2V training

This commit is contained in:
Jaret Burkett
2026-08-03 10:17:39 -06:00
parent 73cab2acf5
commit 8502a845b1
28 changed files with 4095 additions and 41 deletions

View File

@@ -60,6 +60,14 @@ class FileItemDTO(
self.sample_rate = kwargs.get("sample_rate", 48000)
self.num_frames = self.dataset_config.num_frames
self.temporal_compression = kwargs.get("temporal_compression", 8)
# module-level function (picklable) for models whose valid frame
# counts are not temporal_compression * n + 1; None = default math
_sd = kwargs.get("sd", None)
self.frame_count_snapper = (
_sd.get_frame_count_snapper()
if _sd is not None and hasattr(_sd, "get_frame_count_snapper")
else None
)
size_database = kwargs.get("size_database", {})
dataset_root = kwargs.get("dataset_root", None)
self.encode_control_in_text_embeddings = kwargs.get(

View File

@@ -508,15 +508,19 @@ class ImageProcessingDTOMixin:
if self.dataset_config.auto_frame_count:
# allow for any length video here but make sure it is temporally compressable.
vid_length_seconds = total_frames / video_fps
desired_num_frames = int(vid_length_seconds * self.dataset_config.fps)
# make sure it is divisible by temporal_compression
desired_num_frames = desired_num_frames // self.temporal_compression * self.temporal_compression
# TODO, all models currently add a key frame, but future models may not, update here if this changes.
desired_num_frames += 1 # add one for the key frame that is always added
if getattr(self, 'frame_count_snapper', None) is not None:
# model-specific valid-frame-count grid (e.g. minimax_h3's 17n+5)
desired_num_frames = self.frame_count_snapper(desired_num_frames)
else:
# make sure it is divisible by temporal_compression
desired_num_frames = desired_num_frames // self.temporal_compression * self.temporal_compression
# TODO, all models currently add a key frame, but future models may not, update here if this changes.
desired_num_frames += 1 # add one for the key frame that is always added
self.num_frames = desired_num_frames
@@ -705,31 +709,39 @@ class ImageProcessingDTOMixin:
else:
target_duration = source_duration
waveform, sample_rate = torchaudio.load(self.path) # [channels, samples]
waveform = waveform_to_stereo(waveform) # Convert to stereo if not already
if self.dataset_config.audio_normalize:
peak = waveform.abs().amax() # global peak across channels
eps = 1e-9
target_peak = 0.999 # ~ -0.01 dBFS
gain = target_peak / (peak + eps)
waveform = waveform * gain
# torchcodec's AudioDecoder raises when a video has no audio
# track, so probe for a stream before decoding.
import av
with av.open(self.path) as container:
has_audio_stream = len(container.streams.audio) > 0
# Slice to the selected clip region (when we have a meaningful time range)
if source_duration > 0.0:
start_sample = int(round(clip_start_time * sample_rate))
end_sample = int(round(clip_end_time * sample_rate))
start_sample = max(0, min(start_sample, waveform.shape[-1]))
end_sample = max(0, min(end_sample, waveform.shape[-1]))
if end_sample > start_sample:
waveform = waveform[..., start_sample:end_sample]
waveform = None
if has_audio_stream:
waveform, sample_rate = torchaudio.load(self.path) # [channels, samples]
waveform = waveform_to_stereo(waveform) # Convert to stereo if not already
if self.dataset_config.audio_normalize:
peak = waveform.abs().amax() # global peak across channels
eps = 1e-9
target_peak = 0.999 # ~ -0.01 dBFS
gain = target_peak / (peak + eps)
waveform = waveform * gain
# Slice to the selected clip region (when we have a meaningful time range)
if source_duration > 0.0:
start_sample = int(round(clip_start_time * sample_rate))
end_sample = int(round(clip_end_time * sample_rate))
start_sample = max(0, min(start_sample, waveform.shape[-1]))
end_sample = max(0, min(end_sample, waveform.shape[-1]))
if end_sample > start_sample:
waveform = waveform[..., start_sample:end_sample]
else:
# No valid audio segment
waveform = None
else:
# No valid audio segment
# If we can't compute a meaningful time range, treat as no-audio
waveform = None
else:
# If we can't compute a meaningful time range, treat as no-audio
waveform = None
if waveform is not None and waveform.numel() > 0:
target_samples = int(round(target_duration * sample_rate))

View File

@@ -274,12 +274,24 @@ class BaseModel:
except:
# if we have a custom vae, it might not have this
divisibility = 8
# flux packs this again,
if self.is_flux:
divisibility = divisibility * 2
return divisibility
def get_frame_count_snapper(self):
"""Optional hook for video models whose VAE accepts frame counts on a
grid other than the default ``temporal_compression * n + 1``.
Return a MODULE-LEVEL function ``(num_frames: int) -> int`` that snaps
an arbitrary frame count DOWN to the nearest count the video VAE can
encode (it must be picklable by reference — file items travel into
dataloader workers, so no lambdas or bound methods). Returning None
keeps the default auto_frame_count behavior.
"""
return None
# these must be implemented in child classes
def load_model(self):
# override this in child classes

View File

@@ -8,7 +8,7 @@ DIFFUSERS_CONFIGS_ROOT = os.path.join(TOOLKIT_ROOT, "toolkit", "diffusers_config
COMFY_MODELS_PATH = None
# check if ENV variable is set
if 'MODELS_PATH' in os.environ:
if 'MODELS_PATH' in os.environ and os.environ['MODELS_PATH'].strip() != "":
MODELS_PATH = os.environ['MODELS_PATH']
else:
MODELS_PATH = os.path.join(TOOLKIT_ROOT, "models")

View File

@@ -0,0 +1,175 @@
"""Import ComfyUI pre-quantized checkpoints onto toolkit modules.
ComfyUI quantized checkpoints mark each quantized submodule with a
``<prefix>.comfy_quant`` uint8 tensor holding a JSON config, alongside the
quantized ``weight`` and its scale tensors. This module walks those markers
and converts the matching submodules in place:
- ``{"format": "int8_tensorwise", "convrot": true, "convrot_groupsize": G}``
per-output-row symmetric int8 on regular-Hadamard-rotated weights — the
exact storage of the toolkit's convrot8 backend
(toolkit/util/convrot_quant.py:ConvRotInt8Quantizer), so the tensors are
attached to its buffers directly (no requantization). Without the
``convrot`` flag the rotation block is 1, i.e. plain per-row int8, which
the same backend also decodes (rotate is the identity at rot_size 1).
- ``{"format": "nvfp4"}`` block-16 fp4 with e4m3 block scales, an fp32
per-tensor scale and an optional AWQ ``pre_quant_scale`` — attached to
the nvfp4 backend (toolkit/util/nvfp4_quant.py).
- an int8 marker on an ``nn.Embedding`` swaps in :class:`Int8Embedding`
(per-row scales, dequantized per lookup).
Linears become OstrisLinear (class swap in place, like
convert_linear_to_ostris), so LoRA attachment, memory management and the
quantized save paths all work unchanged.
"""
import json
from typing import Dict, Tuple
import torch
from toolkit.util.nvfp4_quant import (
Nvfp4Quantizer,
swap_nvfp4_nibbles,
unswizzle_nvfp4_scales,
)
from toolkit.util.ostris_quant import OstrisLinear, get_ostris_quantizer
def parse_comfy_quant_blob(blob: torch.Tensor) -> dict:
return json.loads(bytes(blob.cpu().tolist()).decode("utf-8"))
class Int8Embedding(torch.nn.Module):
"""An embedding table stored as per-row symmetric int8. Rows are
dequantized per lookup, so the full-precision table never materializes."""
def __init__(self, qweight: torch.Tensor, scales: torch.Tensor, dtype: torch.dtype):
super().__init__()
self.num_embeddings, self.embedding_dim = qweight.shape
self.output_dtype = dtype
self.register_buffer("qweight", qweight.contiguous(), persistent=False)
self.register_buffer(
"scales",
scales.detach().float().reshape(-1).contiguous().view(torch.uint8),
persistent=False,
)
@property
def weight(self):
# full dequantized table, for code that inspects it
scales = self.scales.view(torch.float32)
return (self.qweight.float() * scales.unsqueeze(1)).to(self.output_dtype)
def forward(self, input_ids: torch.Tensor) -> torch.Tensor:
flat = input_ids.reshape(-1)
rows = self.qweight.index_select(0, flat).float()
scales = self.scales.view(torch.float32).index_select(0, flat)
out = rows * scales.unsqueeze(1)
return out.to(self.output_dtype).reshape(*input_ids.shape, self.embedding_dim)
def _to_ostris(module: torch.nn.Linear, quantizer, orig_dtype: torch.dtype) -> OstrisLinear:
if "weight" in module._parameters:
del module._parameters["weight"]
module.ostris_quantizer = quantizer
module.ostris_orig_dtype = orig_dtype
if module.bias is not None:
module.bias.requires_grad_(False)
module.__class__ = OstrisLinear
return module
@torch.no_grad()
def import_comfy_quantized_layers(
root: torch.nn.Module,
state_dict: Dict[str, torch.Tensor],
orig_dtype: torch.dtype = torch.bfloat16,
key_map=None,
) -> Tuple[Dict[str, torch.Tensor], int]:
"""Convert every module a ``comfy_quant`` marker points at and attach its
quantized tensors. Consumes the quantized entries from ``state_dict`` and
returns ``(remaining_state_dict, num_converted)`` — load the remainder
with the regular load_state_dict.
``key_map`` optionally maps a checkpoint prefix to the module path in
``root`` (e.g. comfy text encoder keys onto transformers module paths).
"""
state_dict = dict(state_dict)
converted = 0
marker_keys = [k for k in state_dict.keys() if k.endswith(".comfy_quant")]
for marker_key in marker_keys:
prefix = marker_key[: -len(".comfy_quant")]
conf = parse_comfy_quant_blob(state_dict.pop(marker_key))
fmt = conf.get("format")
module_path = key_map(prefix) if key_map is not None else prefix
module = root.get_submodule(module_path)
weight = state_dict.pop(f"{prefix}.weight")
weight_scale = state_dict.pop(f"{prefix}.weight_scale", None)
if isinstance(module, torch.nn.Embedding):
if fmt != "int8_tensorwise":
raise ValueError(
f"Unsupported comfy quant format {fmt!r} on embedding {prefix}"
)
parent_path, _, attr = module_path.rpartition(".")
parent = root.get_submodule(parent_path) if parent_path else root
setattr(parent, attr, Int8Embedding(weight, weight_scale, orig_dtype))
converted += 1
continue
if not isinstance(module, torch.nn.Linear):
raise ValueError(
f"comfy_quant marker {prefix} points at {type(module).__name__}, "
"expected nn.Linear or nn.Embedding"
)
if fmt == "int8_tensorwise":
rot = int(conf.get("convrot_groupsize", 256)) if conf.get("convrot") else 1
quantizer = get_ostris_quantizer("convrot8")
module.register_buffer("cr8_qdata", weight.contiguous(), persistent=False)
module.register_buffer(
"cr8_scales",
weight_scale.detach().float().reshape(-1).contiguous().view(torch.uint8),
persistent=False,
)
module.cr8_rot_size = rot
elif fmt == "nvfp4":
quantizer = get_ostris_quantizer("nvfp4")
# normalize comfy_kitchen's storage to the toolkit's conventions:
# fp4 pairs are packed high-nibble-first and the e4m3 block scales
# are stored in the swizzled cuBLAS 128x4 tile layout
scales = unswizzle_nvfp4_scales(
weight_scale.view(torch.float8_e4m3fn),
module.out_features,
module.in_features // 16,
)
Nvfp4Quantizer.attach_(
module,
packed=swap_nvfp4_nibbles(weight),
scales=scales,
pts=state_dict.pop(f"{prefix}.weight_scale_2"),
pre_scale=state_dict.pop(f"{prefix}.pre_quant_scale", None),
)
else:
raise ValueError(
f"Unsupported comfy quant format {fmt!r} on {prefix} "
"(supported: int8_tensorwise, nvfp4)"
)
# drop unused calibration extras if present
state_dict.pop(f"{prefix}.input_scale", None)
_to_ostris(module, quantizer, orig_dtype)
bias = state_dict.pop(f"{prefix}.bias", None)
if bias is not None and module.bias is not None:
# bias may still be a meta parameter when the model was built under
# a meta device context
module._parameters["bias"] = torch.nn.Parameter(
bias.detach().clone(), requires_grad=False
)
converted += 1
return state_dict, converted

155
toolkit/util/nvfp4_quant.py Normal file
View File

@@ -0,0 +1,155 @@
"""NVFP4 (ModelOpt/ComfyUI-style) OstrisQuantizer backend — qtype "nvfp4".
Plain block-16 nvfp4 weight storage without ConvRot's rotation: fp4 e2m1
codes packed two per byte, one fp8 e4m3 scale per 16 elements, one fp32
per-tensor scale, plus an optional AWQ ``pre_quant_scale`` applied
elementwise to the input activation before the matmul (ModelOpt convention,
matching ComfyUI's quantized ops).
This is the layout ComfyUI checkpoints tagged ``{"format": "nvfp4"}`` carry
(e.g. the Comfy-Org MiniMax-H3 text encoder). Those exports set
``full_precision_matrix_mult`` — the activations are NOT fp4-quantized — so
the forward here is always the dequantized matmul in the activation's dtype:
weights stay at ~4.25 bits in memory and the math runs on any GPU (or CPU),
no Blackwell fp4 tensor cores required. The triton dequant kernel in
convrot_quant is used when available; pure torch otherwise.
Quantized state attached to each module (uint8 byte views, like the convrot
backends, so nn.Module._apply dtype casts can't corrupt them):
nv4_qdata packed e2m1 codes (uint8, out x in/2; low nibble = even column)
nv4_scales e4m3 block scales (out x in/16)
nv4_pts fp32 per-tensor scale (1 element)
nv4_pre_scale optional fp32 AWQ input scale (in,)
"""
from typing import Optional
import torch
import torch.nn.functional as F
from toolkit.util.convrot_quant import BLOCK, dequantize_nvfp4, quantize_nvfp4
from toolkit.util.ostris_quant import OstrisQuantizer
from toolkit.print import print_acc
NVFP4_QTYPES = ("nvfp4",)
_skip_warned = set()
def unswizzle_nvfp4_scales(scales: torch.Tensor, rows: int, cols: int) -> torch.Tensor:
"""Undo the cuBLAS 128x4-tile block-scale layout (comfy_kitchen's
``to_blocked``) back to a row-major (rows, cols) matrix. ComfyUI nvfp4
checkpoints store ``weight_scale`` swizzled; when the dims are already
tile-aligned the shape is unchanged and only the element order differs."""
n_row_blocks = (rows + 127) // 128
n_col_blocks = (cols + 3) // 4
padded_rows = n_row_blocks * 128
padded_cols = n_col_blocks * 4
x = scales.reshape(-1, 32, 16)
x = x.reshape(-1, 32, 4, 4).transpose(1, 2)
x = x.reshape(n_row_blocks, n_col_blocks, 4, 32, 4)
x = x.reshape(n_row_blocks, n_col_blocks, 128, 4)
x = x.permute(0, 2, 1, 3).reshape(padded_rows, padded_cols)
return x[:rows, :cols].contiguous()
def swap_nvfp4_nibbles(packed: torch.Tensor) -> torch.Tensor:
"""ComfyUI packs fp4 pairs high-nibble-first; the toolkit's decode is
low-nibble-first. Swapping nibbles converts between the two."""
return ((packed << 4) | (packed >> 4)).contiguous()
class Nvfp4Quantizer(OstrisQuantizer):
"""Block-16 nvfp4 weights, full-precision activations. One instance is
shareable across modules."""
def can_quantize(self, module: torch.nn.Linear) -> bool:
if module.in_features % BLOCK != 0:
if module.in_features not in _skip_warned:
_skip_warned.add(module.in_features)
print_acc(
f"nvfp4: skipping linears with in_features={module.in_features} "
f"(needs in divisible by {BLOCK})"
)
return False
return True
def quantize_(self, module: torch.nn.Linear, weight_fp32: torch.Tensor) -> None:
packed, scales, pts = quantize_nvfp4(weight_fp32, optimize_scales=True)
self.attach_(module, packed, scales, pts, pre_scale=None)
@staticmethod
def attach_(
module: torch.nn.Module,
packed: torch.Tensor, # uint8 (out, in/2)
scales: torch.Tensor, # float8_e4m3fn (out, in/16)
pts: torch.Tensor, # fp32 scalar per-tensor scale
pre_scale: Optional[torch.Tensor] = None, # (in,) AWQ input scale
) -> None:
"""Register the quantized representation on the module. Used both by
quantize_ and by importers of pre-quantized checkpoints."""
module.register_buffer("nv4_qdata", packed.contiguous(), persistent=False)
module.register_buffer(
"nv4_scales", scales.contiguous().view(torch.uint8), persistent=False
)
module.register_buffer(
"nv4_pts",
pts.detach().float().clone().reshape(1).view(torch.uint8),
persistent=False,
)
if pre_scale is not None:
module.register_buffer(
"nv4_pre_scale",
pre_scale.detach().float().clone().contiguous().view(torch.uint8),
persistent=False,
)
@staticmethod
def _pts(module) -> torch.Tensor:
return module.nv4_pts.view(torch.float32).reshape(())
@staticmethod
def _pre_scale(module) -> Optional[torch.Tensor]:
buf = getattr(module, "nv4_pre_scale", None)
return None if buf is None else buf.view(torch.float32)
def _dequantize_weight(self, module, dtype: torch.dtype) -> torch.Tensor:
return dequantize_nvfp4(
module.nv4_qdata,
module.nv4_scales.view(torch.float8_e4m3fn),
self._pts(module),
module.out_features,
module.in_features,
dtype,
)
def dequantize(self, module) -> torch.Tensor:
"""The weight as stored. NOTE: with an AWQ pre_quant_scale present the
stored weight expects pre-scaled activations; folding the scale back
(w * pre_scale per column) would reconstruct the original-basis weight
but is deliberately not done here — forward() owns that contract."""
return self._dequantize_weight(module, torch.float32)
def dequantize_folded(self, module) -> torch.Tensor:
"""Weight for raw (un-pre-scaled) activations: the AWQ pre_quant_scale
multiplies the input elementwise, which folds into the weight columns."""
w = self._dequantize_weight(module, torch.float32)
pre_scale = self._pre_scale(module)
if pre_scale is not None:
w = w * pre_scale.unsqueeze(0)
return w
def requantize_(self, module, fp_weight: torch.Tensor) -> None:
w = fp_weight.to(device=module.nv4_qdata.device, dtype=torch.float32)
packed, scales, pts = quantize_nvfp4(w, optimize_scales=True)
module.nv4_qdata = packed
module.nv4_scales = scales.view(torch.uint8)
module.nv4_pts = pts.detach().clone().reshape(1).view(torch.uint8)
def forward(self, module, x: torch.Tensor) -> torch.Tensor:
pre_scale = self._pre_scale(module)
if pre_scale is not None:
x = x * pre_scale.to(dtype=x.dtype)
with torch.no_grad():
w = self._dequantize_weight(module, x.dtype)
return F.linear(x, w, module.bias)

View File

@@ -51,6 +51,13 @@ class OstrisQuantizer:
"""Reconstruct the full weight in the original basis, in float32."""
raise NotImplementedError
def dequantize_folded(self, module: "OstrisLinear") -> torch.Tensor:
"""The full weight with any ACTIVATION-side transform folded in, i.e. a
weight that computes the same output on raw activations. Used when
re-quantizing into a different backend (which won't know about this
backend's activation transforms). Default: same as dequantize."""
return self.dequantize(module)
def requantize_(self, module: "OstrisLinear", fp_weight: torch.Tensor) -> None:
"""Re-quantize in place from a full precision weight in the original basis
(used by the continuous merge/reset method)."""
@@ -179,6 +186,7 @@ def get_ostris_quantizer(qtype: str) -> Optional[OstrisQuantizer]:
"""Resolve a qtype string to a quantizer backend instance, or None if the qtype
does not belong to a custom backend. Add new backends here."""
from toolkit.util.convrot_quant import CONVROT_QTYPES, get_convrot_quantizer
from toolkit.util.nvfp4_quant import NVFP4_QTYPES, Nvfp4Quantizer
from toolkit.util.orbit_quant import ORBIT_QTYPES, OrbitQuantizer
from toolkit.util.orbit_vq_quant import ORBIT_VQ_QTYPES, OrbitVQQuantizer
from toolkit.util.uintx_quant import UINTX_QTYPES, UIntXQuantizer
@@ -190,6 +198,8 @@ def get_ostris_quantizer(qtype: str) -> Optional[OstrisQuantizer]:
quantizer = OrbitVQQuantizer(**ORBIT_VQ_QTYPES[qtype])
elif qtype in CONVROT_QTYPES:
quantizer = get_convrot_quantizer(qtype)
elif qtype in NVFP4_QTYPES:
quantizer = Nvfp4Quantizer()
elif qtype in UINTX_QTYPES:
quantizer = UIntXQuantizer(UINTX_QTYPES[qtype])
if quantizer is not None:
@@ -324,8 +334,33 @@ def convert_linear_to_ostris(
module: torch.nn.Linear, quantizer: OstrisQuantizer
) -> bool:
"""Quantize an nn.Linear in place (class swap). Returns True if the module was
converted (or already was), False if it is not a candidate."""
converted (or already was), False if it is not a candidate.
A module that is ALREADY quantized (e.g. loaded from a pre-quantized
checkpoint) is re-quantized into the requested backend when the qtypes
differ: the weight is dequantized and re-quantized layer by layer, so the
full-precision transient never exceeds one layer's weight. Same qtype is
a no-op (the shipped quantization is kept)."""
if isinstance(module, OstrisLinear):
current_qtype = getattr(module.ostris_quantizer, "qtype", None)
if quantizer.qtype is None or current_qtype == quantizer.qtype:
return True
if not quantizer.can_quantize(module):
return True # keep the existing quantization rather than dropping it
# fold any activation-side transform (e.g. an AWQ pre_quant_scale) into
# the weight so the new backend computes the same function on raw inputs
weight = module.ostris_quantizer.dequantize_folded(module).to(
module.ostris_orig_dtype
)
# backend state lives exclusively in buffers; leftover scalar attrs from
# the old backend are inert
module._buffers.clear()
if quantizer.wants_fp32_weight:
quantizer.quantize_(module, weight.to(torch.float32))
else:
quantizer.quantize_(module, weight)
del weight
module.ostris_quantizer = quantizer
return True
weight = getattr(module, "weight", None)
if not isinstance(weight, torch.nn.Parameter) or not weight.dtype.is_floating_point:

View File

@@ -176,7 +176,16 @@ def quantize(
try:
# check if m is QLinear or QConv2d
if m.__class__.__name__ in Q_MODULES:
continue
# OstrisLinear may still be RE-quantized into a different
# ostris qtype (same qtype is a per-layer no-op); every other
# already-quantized module type is always left alone, which
# also keeps quanto/torchao from double-quantizing
# pre-quantized checkpoints
if not (
isinstance(weights, ostristype)
and m.__class__.__name__ == "OstrisLinear"
):
continue
if (
isinstance(weights, aotype)
and not isinstance(m, torch.nn.Linear)
@@ -193,7 +202,10 @@ def quantize(
continue
orig_device = None
if quantize_device is not None and next(m.children(), None) is None:
# OstrisLinear layers being re-quantized hold buffers, not params
param = next(m.parameters(recurse=False), None)
if param is None:
param = next(m.buffers(recurse=False), None)
if param is not None:
orig_device = param.device
m.to(quantize_device)