Text-to-Image
Diffusers
Safetensors
MageFlowPipeline
ajh
mage-flow
mage-flow-turbo
mage-flow-turbo-nvfp4-balanced-ajh
nvfp4
fp8
blackwell
qwen3-vl
quantization
balanced
Instructions to use ajh-code/Mage-Flow-Turbo-NVFP4-Balanced-AJH with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use ajh-code/Mage-Flow-Turbo-NVFP4-Balanced-AJH with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("ajh-code/Mage-Flow-Turbo-NVFP4-Balanced-AJH", torch_dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- Draw Things
- DiffusionBee
| """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, | |
| }, | |
| ) | |