Mage-Flow-NVFP4-Balanced-AJH / runtime /quant_text_encoder.py
ajh-code's picture
Add files using upload-large-folder tool
ca8b3a9 verified
Raw
History Blame Contribute Delete
13.5 kB
"""Load the packaged mixed NVFP4/FP8 Qwen3-VL text encoder without BF16 shards."""
from __future__ import annotations
from collections import Counter
import hashlib
import json
from pathlib import Path
from typing import Any
import torch
import torch.nn.functional as F
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
from safetensors import safe_open
from transformers import AutoConfig, AutoProcessor, AutoTokenizer
EXPECTED_ARTIFACT_SHA256 = (
"719906b435800757d22013d3d475a4853d59b779b022669fa8b8a193b85d0f41"
)
EXPECTED_FORMAT_COUNTS = {"nvfp4": 224, "float8_e4m3fn": 14}
EXPECTED_PROJECTIONS = {
"mlp.down_proj",
"mlp.gate_proj",
"mlp.up_proj",
"self_attn.k_proj",
"self_attn.o_proj",
"self_attn.q_proj",
"self_attn.v_proj",
}
def sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1 << 20), b""):
digest.update(chunk)
return digest.hexdigest()
def _decode_quant_config(tensor: torch.Tensor) -> dict[str, Any]:
payload = bytes(tensor.cpu().to(torch.uint8).tolist()).decode("utf-8")
parsed = json.loads(payload)
if not isinstance(parsed, dict) or not isinstance(parsed.get("format"), str):
raise ValueError(f"invalid comfy_quant payload: {payload!r}")
return parsed
def _artifact_layer_to_hf_module(layer_key: str) -> str:
prefix = "model.layers."
if not layer_key.startswith(prefix):
raise ValueError(f"quantized layer is outside the language stack: {layer_key}")
remainder = layer_key[len(prefix) :]
layer_text, projection = remainder.split(".", 1)
layer_index = int(layer_text)
if layer_index < 0 or layer_index >= 36:
raise ValueError(f"language layer index is out of range: {layer_key}")
if projection not in EXPECTED_PROJECTIONS:
raise ValueError(f"unexpected quantized projection: {layer_key}")
return f"model.language_model.layers.{layer_index}.{projection}"
def _artifact_weight_to_hf_name(key: str) -> str | None:
if key == "model.embed_tokens.weight":
return "model.language_model.embed_tokens.weight"
if key == "model.norm.weight":
return "model.language_model.norm.weight"
if key.startswith("model.layers."):
return "model.language_model.layers." + key.removeprefix("model.layers.")
if key.startswith("model.visual."):
return key
return None
def _resolve_parent(module: torch.nn.Module, dotted_name: str) -> tuple[Any, str]:
parts = dotted_name.split(".")
parent: Any = module
for part in parts[:-1]:
parent = getattr(parent, part)
return parent, parts[-1]
class PublishedQuantLinear(torch.nn.Module):
"""Inference-only projection backed by comfy-kitchen packed tensors."""
def __init__(
self,
*,
in_features: int,
out_features: int,
quant_format: str,
qdata: torch.Tensor,
weight_scale: torch.Tensor,
weight_scale_2: torch.Tensor | None,
) -> None:
super().__init__()
self.in_features = int(in_features)
self.out_features = int(out_features)
self.quant_format = str(quant_format)
self.register_buffer("qdata", qdata.clone().contiguous(), persistent=True)
self.register_buffer(
"weight_scale",
weight_scale.clone().contiguous(),
persistent=True,
)
if weight_scale_2 is None:
self.weight_scale_2 = None
else:
self.register_buffer(
"weight_scale_2",
weight_scale_2.clone().contiguous(),
persistent=True,
)
def _weight_quantized_tensor(self) -> Any:
from comfy_kitchen.tensor import (
QuantizedTensor,
TensorCoreFP8Layout,
TensorCoreNVFP4Layout,
)
shape = (self.out_features, self.in_features)
if self.quant_format == "nvfp4":
if self.weight_scale_2 is None:
raise RuntimeError("NVFP4 projection is missing its second scale")
params = TensorCoreNVFP4Layout.Params(
scale=self.weight_scale_2,
orig_dtype=torch.bfloat16,
orig_shape=shape,
block_scale=self.weight_scale,
)
return QuantizedTensor(
self.qdata,
"TensorCoreNVFP4Layout",
params,
)
if self.quant_format == "float8_e4m3fn":
params = TensorCoreFP8Layout.Params(
scale=self.weight_scale,
orig_dtype=torch.bfloat16,
orig_shape=shape,
)
return QuantizedTensor(
self.qdata,
"TensorCoreFP8Layout",
params,
)
raise ValueError(f"unsupported quantized format: {self.quant_format}")
def forward(self, value: torch.Tensor) -> torch.Tensor:
from comfy_kitchen.tensor import QuantizedTensor
input_shape = tuple(value.shape)
flattened = value.reshape(-1, input_shape[-1]).contiguous()
layout = (
"TensorCoreNVFP4Layout"
if self.quant_format == "nvfp4"
else "TensorCoreFP8Layout"
)
input_quantized = QuantizedTensor.from_float(flattened, layout)
output = F.linear(
input_quantized,
self._weight_quantized_tensor(),
None,
)
return output.reshape(*input_shape[:-1], self.out_features)
def _install_quantized_linears(
hf_module: torch.nn.Module,
artifact_path: Path,
) -> dict[str, Any]:
format_counts: Counter[str] = Counter()
installed: list[str] = []
storage_bytes = 0
original_bf16_bytes = 0
with safe_open(str(artifact_path), framework="pt", device="cpu") as handle:
config_keys = sorted(
key for key in handle.keys() if key.endswith(".comfy_quant")
)
for config_key in config_keys:
layer_key = config_key.removesuffix(".comfy_quant")
quant_format = _decode_quant_config(
handle.get_tensor(config_key)
)["format"]
module_name = _artifact_layer_to_hf_module(layer_key)
parent, leaf = _resolve_parent(hf_module, module_name)
original = getattr(parent, leaf)
if not isinstance(original, torch.nn.Linear):
raise TypeError(
f"{module_name}: expected torch.nn.Linear, got "
f"{type(original).__name__}"
)
if original.bias is not None:
raise ValueError(f"{module_name}: quantized projection has a bias")
qdata = handle.get_tensor(f"{layer_key}.weight")
weight_scale = handle.get_tensor(f"{layer_key}.weight_scale")
weight_scale_2 = (
handle.get_tensor(f"{layer_key}.weight_scale_2")
if quant_format == "nvfp4"
else None
)
replacement = PublishedQuantLinear(
in_features=original.in_features,
out_features=original.out_features,
quant_format=quant_format,
qdata=qdata,
weight_scale=weight_scale,
weight_scale_2=weight_scale_2,
)
setattr(parent, leaf, replacement)
format_counts[quant_format] += 1
installed.append(module_name)
original_bf16_bytes += (
original.in_features * original.out_features * 2
)
storage_bytes += replacement.qdata.nbytes
storage_bytes += replacement.weight_scale.nbytes
if replacement.weight_scale_2 is not None:
storage_bytes += replacement.weight_scale_2.nbytes
summary = {
"installed_module_count": len(installed),
"format_counts": dict(format_counts),
"original_projection_bf16_bytes": original_bf16_bytes,
"packed_projection_bytes": storage_bytes,
"projection_saving_bytes": original_bf16_bytes - storage_bytes,
"projection_saving_gib": (
(original_bf16_bytes - storage_bytes) / float(1 << 30)
),
}
if (
summary["installed_module_count"] != 238
or summary["format_counts"] != EXPECTED_FORMAT_COUNTS
):
raise RuntimeError(f"unexpected text quantization policy: {summary}")
return summary
def _load_nonquantized_weights(
hf_module: torch.nn.Module,
artifact_path: Path,
) -> dict[str, Any]:
loaded: list[str] = []
unexpected: list[str] = []
with safe_open(str(artifact_path), framework="pt", device="cpu") as handle:
for artifact_key in sorted(handle.keys()):
if (
artifact_key.endswith(".comfy_quant")
or artifact_key.endswith(".weight_scale")
or artifact_key.endswith(".weight_scale_2")
):
continue
target_name = _artifact_weight_to_hf_name(artifact_key)
if target_name is None:
unexpected.append(artifact_key)
continue
try:
parent, leaf = _resolve_parent(hf_module, target_name)
except AttributeError:
unexpected.append(artifact_key)
continue
current = getattr(parent, leaf, None)
if isinstance(current, PublishedQuantLinear):
continue
if target_name.endswith(".weight"):
projection_name = target_name.removesuffix(".weight")
try:
projection_parent, projection_leaf = _resolve_parent(
hf_module, projection_name
)
if isinstance(
getattr(projection_parent, projection_leaf),
PublishedQuantLinear,
):
continue
except AttributeError:
pass
set_module_tensor_to_device(
hf_module,
target_name,
"cpu",
value=handle.get_tensor(artifact_key),
)
loaded.append(target_name)
hf_module.tie_weights()
meta_parameters = [
name for name, value in hf_module.named_parameters() if value.is_meta
]
meta_buffers = [
name for name, value in hf_module.named_buffers() if value.is_meta
]
if meta_parameters or meta_buffers:
raise RuntimeError(
"packed text loader left unresolved meta tensors: "
f"{(meta_parameters + meta_buffers)[:4]}"
)
if unexpected:
raise RuntimeError(
f"packed text artifact contains unmapped tensors: {unexpected[:4]}"
)
return {
"loaded_nonquantized_tensor_count": len(loaded),
"unresolved_meta_parameters": meta_parameters,
"unresolved_meta_buffers": meta_buffers,
}
def load_quantized_text_encoder(
*,
text_encoder_dir: str | Path,
artifact_path: str | Path,
tokenizer_max_length: int,
dit_structure: dict[str, Any],
use_packed_text_infer: bool,
) -> tuple[torch.nn.Module, dict[str, Any]]:
"""Construct Mage's text wrapper directly from the packaged quant artifact."""
from mage_flow.models.modules.text_encoder import (
CustomQwen3VLForConditionalGeneration,
TextEncoder,
)
text_encoder_dir = Path(text_encoder_dir).resolve()
artifact_path = Path(artifact_path).resolve()
if sha256(artifact_path) != EXPECTED_ARTIFACT_SHA256:
raise RuntimeError("packaged text-encoder artifact SHA-256 mismatch")
config = AutoConfig.from_pretrained(
str(text_encoder_dir),
local_files_only=True,
)
with init_empty_weights():
hf_module = CustomQwen3VLForConditionalGeneration._from_config(
config,
attn_implementation="flash_attention_2",
dtype=torch.bfloat16,
)
quant_summary = _install_quantized_linears(hf_module, artifact_path)
load_summary = _load_nonquantized_weights(hf_module, artifact_path)
text_encoder = TextEncoder.__new__(TextEncoder)
torch.nn.Module.__init__(text_encoder)
text_encoder.model_name = str(text_encoder_dir)
text_encoder.tokenizer_max_length = int(tokenizer_max_length)
text_encoder.tokenizer = AutoTokenizer.from_pretrained(
str(text_encoder_dir),
local_files_only=True,
)
text_encoder.tokenizer.padding_side = "right"
text_encoder.processor = AutoProcessor.from_pretrained(
str(text_encoder_dir),
local_files_only=True,
)
text_encoder.hf_module = hf_module.eval().requires_grad_(False)
text_encoder.prompt_template_encode = ""
text_encoder.prompt_template_encode_start_idx = 0
text_encoder.dit_structure = dict(dit_structure)
text_encoder.use_packed_text_infer = bool(use_packed_text_infer)
text_encoder.eval().requires_grad_(False)
return (
text_encoder,
{
"artifact": str(artifact_path),
"artifact_sha256": EXPECTED_ARTIFACT_SHA256,
"quantized": quant_summary,
"nonquantized": load_summary,
},
)