[diffusion] feat: automatically infer comfy fp8 activation scaling (#36060)
This commit is contained in:
@@ -273,8 +273,10 @@ sglang serve \
|
|||||||
--port 30010
|
--port 30010
|
||||||
```
|
```
|
||||||
|
|
||||||
SGLang uses its native static-activation FP8 linear path for attention and
|
SGLang uses its native static-activation FP8 linear path when the checkpoint
|
||||||
`fc1`. The checkpoint marks `fc2` for full-precision matrix multiplication, so
|
stores input scales, and automatically uses dynamic activation scaling for
|
||||||
|
Comfy FP8 exports that omit them. The checkpoint above marks `fc2` for
|
||||||
|
full-precision matrix multiplication, so
|
||||||
SGLang retains its FP8 storage but materializes and scales one compute-dtype
|
SGLang retains its FP8 storage but materializes and scales one compute-dtype
|
||||||
`fc2` matrix for each call. This preserves the checkpoint's mixed execution
|
`fc2` matrix for each call. This preserves the checkpoint's mixed execution
|
||||||
contract and low resident weight memory, but that part is slower than a fully
|
contract and low resident weight memory, but that part is slower than a fully
|
||||||
|
|||||||
@@ -51,7 +51,9 @@ repo contains multiple candidate checkpoints, pass
|
|||||||
`--transformer-weights-path` explicitly.
|
`--transformer-weights-path` explicitly.
|
||||||
|
|
||||||
MiniMax-H3 is a verified example for Comfy safetensors with per-layer metadata,
|
MiniMax-H3 is a verified example for Comfy safetensors with per-layer metadata,
|
||||||
including `pruned_fp8_scaled` and serialized ConvRot INT8. Pass one selected
|
including `pruned_fp8_scaled` and serialized ConvRot INT8. Other Comfy FP8
|
||||||
|
exports are also auto-detected: the presence of an input scale selects static
|
||||||
|
activation scaling, while its absence selects dynamic scaling. Pass one selected
|
||||||
FL2VA or Ref2VA DiT file by local path, `owner/repo/path/file.safetensors`, or
|
FL2VA or Ref2VA DiT file by local path, `owner/repo/path/file.safetensors`, or
|
||||||
direct Hugging Face file URL; do not combine it with `--quantization`. Its GGUF
|
direct Hugging Face file URL; do not combine it with `--quantization`. Its GGUF
|
||||||
usage is documented in
|
usage is documented in
|
||||||
|
|||||||
@@ -94,10 +94,13 @@ class ComfyFp8Config(QuantizationConfig):
|
|||||||
super().__init__()
|
super().__init__()
|
||||||
self.layer_markers = layer_markers
|
self.layer_markers = layer_markers
|
||||||
self.selected: list[str] = []
|
self.selected: list[str] = []
|
||||||
self._fp8_config = Fp8Config(
|
self._fp8_configs = {
|
||||||
is_checkpoint_fp8_serialized=True,
|
activation_scheme: Fp8Config(
|
||||||
activation_scheme="static",
|
is_checkpoint_fp8_serialized=True,
|
||||||
)
|
activation_scheme=activation_scheme,
|
||||||
|
)
|
||||||
|
for activation_scheme in ("static", "dynamic")
|
||||||
|
}
|
||||||
|
|
||||||
unsupported = {
|
unsupported = {
|
||||||
prefix: marker.get("format")
|
prefix: marker.get("format")
|
||||||
@@ -140,7 +143,8 @@ class ComfyFp8Config(QuantizationConfig):
|
|||||||
self.selected.append(prefix)
|
self.selected.append(prefix)
|
||||||
if marker.get("full_precision_matrix_mult", False):
|
if marker.get("full_precision_matrix_mult", False):
|
||||||
return ComfyFullPrecisionFp8LinearMethod()
|
return ComfyFullPrecisionFp8LinearMethod()
|
||||||
return Fp8LinearMethod(self._fp8_config)
|
activation_scheme = marker.get("_activation_scheme", "static")
|
||||||
|
return Fp8LinearMethod(self._fp8_configs[activation_scheme])
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
|||||||
@@ -77,10 +77,6 @@ def inspect_comfy_quant_markers(
|
|||||||
for prefix, marker in raw_markers.items():
|
for prefix, marker in raw_markers.items():
|
||||||
marker_format = marker.get("format")
|
marker_format = marker.get("format")
|
||||||
required = {f"{prefix}.weight", f"{prefix}.weight_scale"}
|
required = {f"{prefix}.weight", f"{prefix}.weight_scale"}
|
||||||
if marker_format == "float8_e4m3fn" and not marker.get(
|
|
||||||
"full_precision_matrix_mult", False
|
|
||||||
):
|
|
||||||
required.add(f"{prefix}.input_scale")
|
|
||||||
if marker_format not in ("float8_e4m3fn", "int8_tensorwise"):
|
if marker_format not in ("float8_e4m3fn", "int8_tensorwise"):
|
||||||
continue
|
continue
|
||||||
missing = required - checkpoint_meta.keys()
|
missing = required - checkpoint_meta.keys()
|
||||||
@@ -89,7 +85,10 @@ def inspect_comfy_quant_markers(
|
|||||||
f"Comfy layer {prefix!r} is missing checkpoint tensors: "
|
f"Comfy layer {prefix!r} is missing checkpoint tensors: "
|
||||||
f"{sorted(missing)}"
|
f"{sorted(missing)}"
|
||||||
)
|
)
|
||||||
if marker_format != "int8_tensorwise":
|
if marker_format == "float8_e4m3fn":
|
||||||
|
marker["_activation_scheme"] = (
|
||||||
|
"static" if f"{prefix}.input_scale" in checkpoint_meta else "dynamic"
|
||||||
|
)
|
||||||
continue
|
continue
|
||||||
weight_dtype, weight_shape = checkpoint_meta[f"{prefix}.weight"]
|
weight_dtype, weight_shape = checkpoint_meta[f"{prefix}.weight"]
|
||||||
scale_dtype, scale_shape = checkpoint_meta[f"{prefix}.weight_scale"]
|
scale_dtype, scale_shape = checkpoint_meta[f"{prefix}.weight_scale"]
|
||||||
|
|||||||
@@ -275,7 +275,7 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
self.assertEqual(curve_shape, (1025, 8))
|
self.assertEqual(curve_shape, (1025, 8))
|
||||||
self.assertEqual(comfy_quant["blocks.0.mlp.fc1"]["format"], "int8_tensorwise")
|
self.assertEqual(comfy_quant["blocks.0.mlp.fc1"]["format"], "int8_tensorwise")
|
||||||
|
|
||||||
def test_inspect_minimax_h3_fp8_validates_required_scales(self):
|
def test_inspect_minimax_h3_fp8_detects_static_activation_scale(self):
|
||||||
marker = torch.tensor(list(b'{"format":"float8_e4m3fn"}'), dtype=torch.uint8)
|
marker = torch.tensor(list(b'{"format":"float8_e4m3fn"}'), dtype=torch.uint8)
|
||||||
with tempfile.NamedTemporaryFile(suffix=".safetensors") as f:
|
with tempfile.NamedTemporaryFile(suffix=".safetensors") as f:
|
||||||
save_file(
|
save_file(
|
||||||
@@ -292,7 +292,31 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
|
|
||||||
_, layer_markers = inspect_minimax_h3_safetensors([f.name])
|
_, layer_markers = inspect_minimax_h3_safetensors([f.name])
|
||||||
|
|
||||||
self.assertEqual(layer_markers["blocks.0.mlp.fc1"], {"format": "float8_e4m3fn"})
|
self.assertEqual(
|
||||||
|
layer_markers["blocks.0.mlp.fc1"],
|
||||||
|
{"format": "float8_e4m3fn", "_activation_scheme": "static"},
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_inspect_minimax_h3_fp8_without_input_scale_uses_dynamic_activation(self):
|
||||||
|
marker = torch.tensor(list(b'{"format":"float8_e4m3fn"}'), dtype=torch.uint8)
|
||||||
|
with tempfile.NamedTemporaryFile(suffix=".safetensors") as f:
|
||||||
|
save_file(
|
||||||
|
{
|
||||||
|
"blocks.0.mlp.fc1.weight": torch.ones(
|
||||||
|
(2, 2), dtype=torch.float8_e4m3fn
|
||||||
|
),
|
||||||
|
"blocks.0.mlp.fc1.weight_scale": torch.tensor(0.5),
|
||||||
|
"blocks.0.mlp.fc1.comfy_quant": marker,
|
||||||
|
},
|
||||||
|
f.name,
|
||||||
|
)
|
||||||
|
|
||||||
|
_, layer_markers = inspect_minimax_h3_safetensors([f.name])
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
layer_markers["blocks.0.mlp.fc1"],
|
||||||
|
{"format": "float8_e4m3fn", "_activation_scheme": "dynamic"},
|
||||||
|
)
|
||||||
|
|
||||||
def test_minimax_h3_comfy_int8_resolves_serialized_kitchen(self):
|
def test_minimax_h3_comfy_int8_resolves_serialized_kitchen(self):
|
||||||
config = resolve_minimax_h3_checkpoint_quantization(
|
config = resolve_minimax_h3_checkpoint_quantization(
|
||||||
@@ -357,7 +381,10 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
def test_minimax_h3_comfy_fp8_resolves_per_layer_dispatch(self):
|
def test_minimax_h3_comfy_fp8_resolves_per_layer_dispatch(self):
|
||||||
config = resolve_minimax_h3_checkpoint_quantization(
|
config = resolve_minimax_h3_checkpoint_quantization(
|
||||||
{
|
{
|
||||||
"blocks.0.attn.qkv_proj": {"format": "float8_e4m3fn"},
|
"blocks.0.attn.qkv_proj": {
|
||||||
|
"format": "float8_e4m3fn",
|
||||||
|
"_activation_scheme": "dynamic",
|
||||||
|
},
|
||||||
"blocks.0.mlp.fc2": {
|
"blocks.0.mlp.fc2": {
|
||||||
"format": "float8_e4m3fn",
|
"format": "float8_e4m3fn",
|
||||||
"full_precision_matrix_mult": True,
|
"full_precision_matrix_mult": True,
|
||||||
@@ -372,10 +399,9 @@ class TestTransformerQuantHelpers(unittest.TestCase):
|
|||||||
config.get_quant_method(layer, "blocks.0.mlp.fc2"),
|
config.get_quant_method(layer, "blocks.0.mlp.fc2"),
|
||||||
ComfyFullPrecisionFp8LinearMethod,
|
ComfyFullPrecisionFp8LinearMethod,
|
||||||
)
|
)
|
||||||
self.assertIsInstance(
|
fp8_method = config.get_quant_method(layer, "blocks.0.attn.qkv_proj")
|
||||||
config.get_quant_method(layer, "blocks.0.attn.qkv_proj"),
|
self.assertIsInstance(fp8_method, Fp8LinearMethod)
|
||||||
Fp8LinearMethod,
|
self.assertEqual(fp8_method.quant_config.activation_scheme, "dynamic")
|
||||||
)
|
|
||||||
self.assertIsInstance(
|
self.assertIsInstance(
|
||||||
config.get_quant_method(layer, "unmarked"),
|
config.get_quant_method(layer, "unmarked"),
|
||||||
UnquantizedLinearMethod,
|
UnquantizedLinearMethod,
|
||||||
|
|||||||
Reference in New Issue
Block a user