Add support for MiniMax H3 T2V and I2V training
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
175
toolkit/util/comfy_quant_import.py
Normal file
175
toolkit/util/comfy_quant_import.py
Normal 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
155
toolkit/util/nvfp4_quant.py
Normal 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)
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user