From 2d1f34f9c7e91ee8a67eb12957fb834f325f0d79 Mon Sep 17 00:00:00 2001 From: Haoming Date: Fri, 3 Apr 2026 17:26:57 +0700 Subject: [PATCH] mxfp8 --- backend/float.py | 34 +++++++++++++++++++++++++++ backend/memory_management.py | 14 +++++++++++ backend/operations.py | 3 +++ backend/operations_mixed_precision.py | 12 ++++++++++ backend/quant_ops.py | 33 ++++++++++++++++++++++++++ 5 files changed, 96 insertions(+) diff --git a/backend/float.py b/backend/float.py index 9526a571..30d45f51 100644 --- a/backend/float.py +++ b/backend/float.py @@ -173,3 +173,37 @@ def stochastic_round_quantize_nvfp4_by_block(x, per_tensor_scale, pad_16x, seed= output_block[i : i + slice_size].copy_(block) return output_fp4, to_blocked(output_block, flatten=False) + + +def stochastic_round_quantize_mxfp8_by_block(x, pad_32x, seed=0): + def roundup(x_val, multiple): + return ((x_val + multiple - 1) // multiple) * multiple + + if pad_32x: + rows, cols = x.shape + padded_rows = roundup(rows, 32) + padded_cols = roundup(cols, 32) + if padded_rows != rows or padded_cols != cols: + x = torch.nn.functional.pad(x, (0, padded_cols - cols, 0, padded_rows - rows)) + + F8_E4M3_MAX = 448.0 + E8M0_BIAS = 127 + BLOCK_SIZE = 32 + + rows, cols = x.shape + x_blocked = x.reshape(rows, -1, BLOCK_SIZE) + max_abs = torch.amax(torch.abs(x_blocked), dim=-1) + + scale_needed = torch.clamp(max_abs.float() / F8_E4M3_MAX, min=2 ** (-127)) + exp_biased = torch.clamp(torch.ceil(torch.log2(scale_needed)).to(torch.int32) + E8M0_BIAS, 0, 254) + block_scales_e8m0 = exp_biased.to(torch.uint8) + + zero_mask = max_abs == 0 + block_scales_f32 = (block_scales_e8m0.to(torch.int32) << 23).view(torch.float32) + block_scales_f32 = torch.where(zero_mask, torch.ones_like(block_scales_f32), block_scales_f32) + + data_scaled = (x_blocked.float() / block_scales_f32.unsqueeze(-1)).reshape(rows, cols) + output_fp8 = stochastic_rounding(data_scaled, torch.float8_e4m3fn, seed=seed) + + block_scales_e8m0 = torch.where(zero_mask, torch.zeros_like(block_scales_e8m0), block_scales_e8m0) + return output_fp8, to_blocked(block_scales_e8m0, flatten=False).view(torch.float8_e8m0fnu) diff --git a/backend/memory_management.py b/backend/memory_management.py index ca06ff36..ed66938c 100644 --- a/backend/memory_management.py +++ b/backend/memory_management.py @@ -1295,6 +1295,20 @@ def supports_nvfp4_compute(device: torch.device = None) -> bool: return True +def supports_mxfp8_compute(device: torch.device = None) -> bool: + if not is_nvidia(): + return False + + if torch_version_numeric < (2, 10): + return False + + props = torch.cuda.get_device_properties(device) + if props.major < 10: + return False + + return True + + def extended_fp16_support() -> bool: return torch_version_numeric >= (2, 7) diff --git a/backend/operations.py b/backend/operations.py index dc4ca6a9..f8820066 100644 --- a/backend/operations.py +++ b/backend/operations.py @@ -753,10 +753,13 @@ def using_forge_operations(operations=None, device=None, dtype=None, manual_cast _dtype = torch.bfloat16 if memory_management.should_use_bf16(_device) else torch.float32 fp8_compute = memory_management.supports_fp8_compute(_device) nvfp4_compute = memory_management.supports_nvfp4_compute(_device) + mxfp8_compute = memory_management.supports_mxfp8_compute(_device) disabled = set() if not nvfp4_compute: disabled.add("nvfp4") + if not mxfp8_compute: + disabled.add("mxfp8") if not fp8_compute: disabled.add("float8_e4m3fn") disabled.add("float8_e5m2") diff --git a/backend/operations_mixed_precision.py b/backend/operations_mixed_precision.py index 7876795e..d1bdf5da 100644 --- a/backend/operations_mixed_precision.py +++ b/backend/operations_mixed_precision.py @@ -101,7 +101,19 @@ def mixed_precision_ops(quant_config={}, compute_dtype=torch.bfloat16, full_prec orig_dtype=MixedPrecisionOps._compute_dtype, orig_shape=(self.out_features, self.in_features), ) + elif self.quant_format == "mxfp8": + block_scale = self._load_scale_param(state_dict, prefix, "weight_scale", device, manually_loaded_keys, dtype=torch.uint8) + if block_scale is None: + raise ValueError(f"Missing MXFP8 block scales for layer {layer_name}") + + block_scale = block_scale.view(torch.float8_e8m0fnu) + + params = layout_cls.Params( + scale=block_scale, + orig_dtype=MixedPrecisionOps._compute_dtype, + orig_shape=(self.out_features, self.in_features), + ) elif self.quant_format == "nvfp4": tensor_scale = self._load_scale_param(state_dict, prefix, "weight_scale_2", device, manually_loaded_keys) block_scale = self._load_scale_param(state_dict, prefix, "weight_scale", device, manually_loaded_keys, dtype=torch.float8_e4m3fn) diff --git a/backend/quant_ops.py b/backend/quant_ops.py index 59a44fd5..d29f25c4 100644 --- a/backend/quant_ops.py +++ b/backend/quant_ops.py @@ -6,6 +6,7 @@ from comfy_kitchen.tensor import ( # noqa QuantizedLayout, QuantizedTensor, TensorCoreFP8Layout, + TensorCoreMXFP8Layout, TensorCoreNVFP4Layout, get_layout_class, register_layout_class, @@ -97,6 +98,31 @@ class TensorCoreNVFP4Layout(TensorCoreNVFP4Layout): return qdata, params +class TensorCoreMXFP8Layout(TensorCoreMXFP8Layout): + @classmethod + def quantize(cls, tensor, scale=None, stochastic_rounding=0, inplace_ops=False): + if tensor.dim() != 2: + raise ValueError(f"MXFP8 requires 2D tensor, got {tensor.dim()}D") + + orig_dtype = tensor.dtype + orig_shape = tuple(tensor.shape) + + padded_shape = cls.get_padded_shape(orig_shape) + needs_padding = padded_shape != orig_shape + + if stochastic_rounding > 0: + qdata, block_scale = float.stochastic_round_quantize_mxfp8_by_block(tensor, pad_32x=needs_padding, seed=stochastic_rounding) + else: + qdata, block_scale = ck.quantize_mxfp8(tensor, pad_32x=needs_padding) + + params = cls.Params( + scale=block_scale, + orig_dtype=orig_dtype, + orig_shape=orig_shape, + ) + return qdata, params + + class TensorCoreFP8E4M3Layout(_TensorCoreFP8LayoutBase): FP8_DTYPE = torch.float8_e4m3fn @@ -115,6 +141,7 @@ register_layout_class("TensorCoreFP8Layout", TensorCoreFP8Layout) register_layout_class("TensorCoreFP8E4M3Layout", TensorCoreFP8E4M3Layout) register_layout_class("TensorCoreFP8E5M2Layout", TensorCoreFP8E5M2Layout) register_layout_class("TensorCoreNVFP4Layout", TensorCoreNVFP4Layout) +register_layout_class("TensorCoreMXFP8Layout", TensorCoreMXFP8Layout) QUANT_ALGOS = { "float8_e4m3fn": { @@ -133,4 +160,10 @@ QUANT_ALGOS = { "comfy_tensor_layout": "TensorCoreNVFP4Layout", "group_size": 16, }, + "mxfp8": { + "storage_t": torch.float8_e4m3fn, + "parameters": {"weight_scale", "input_scale"}, + "comfy_tensor_layout": "TensorCoreMXFP8Layout", + "group_size": 32, + }, }