Match Comfy Qwen FP32 projections
This commit is contained in:
parent
e43e9d9d81
commit
ab10de7521
5 changed files with 136 additions and 4 deletions
26
PARITY.md
26
PARITY.md
|
|
@ -149,6 +149,32 @@ investigation must capture/compare the effective `QuantizedTensor` parameters
|
||||||
and direct packed-weight wrapper at dispatch time. Do not revisit attention,
|
and direct packed-weight wrapper at dispatch time. Do not revisit attention,
|
||||||
sampler, or MLP ordering until that exact module-dispatch difference is found.
|
sampler, or MLP ordering until that exact module-dispatch difference is found.
|
||||||
|
|
||||||
|
### Effective Weight Root Cause And Resolution
|
||||||
|
|
||||||
|
The storage-safe dispatch probe found the active difference. Comfy's text
|
||||||
|
encoder calls Qwen with `dtype=torch.float32`; its full-precision NVFP4 branch
|
||||||
|
casts each projection to an effective FP32 matrix before `F.linear`. Direct had
|
||||||
|
constructed the Qwen encoder with BF16, causing its full-precision NVFP4 branch
|
||||||
|
to dequantize the same packed data into BF16.
|
||||||
|
|
||||||
|
The probe proved that packed qdata, tensor scale, and block scale are exact for
|
||||||
|
all Q/K/V/gate/up/down projections. With direct Qwen switched to FP32:
|
||||||
|
|
||||||
|
| Projection check | Result |
|
||||||
|
| --- | --- |
|
||||||
|
| Effective Q/K/V/gate/up/down FP32 weights | Bit-exact to Comfy |
|
||||||
|
| Q/K/V/gate/up linear outputs | Bit-exact to Comfy |
|
||||||
|
| Layer-0 O-projected attention | Mean/max `0` / `0` |
|
||||||
|
| Layer-0 post-attention residual | Mean/max `0` / `0` |
|
||||||
|
| Layer-0 post-attention RMSNorm | Mean/max `0` / `0` |
|
||||||
|
| Layer-0 MLP output | Mean/max `0` / `0` |
|
||||||
|
| Layer-0 final output | Mean/max `0` / `0` |
|
||||||
|
|
||||||
|
The special `down` dispatch sub-probe that supplied Comfy's already-activated
|
||||||
|
input directly to `F.linear` is intentionally invalid because it bypasses the
|
||||||
|
module's required `pre_quant_scale`; the full layer replay above is the valid
|
||||||
|
gate and is exact.
|
||||||
|
|
||||||
## Current Runtime Scope
|
## Current Runtime Scope
|
||||||
|
|
||||||
| Component | Implemented | Known limitation |
|
| Component | Implemented | Known limitation |
|
||||||
|
|
|
||||||
|
|
@ -67,8 +67,8 @@ class Nvfp4Linear(nn.Module):
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
if x.shape[-1] != self.in_features:
|
if x.shape[-1] != self.in_features:
|
||||||
raise ValueError(f"Expected feature width {self.in_features}, received {x.shape[-1]}.")
|
raise ValueError(f"Expected feature width {self.in_features}, received {x.shape[-1]}.")
|
||||||
if x.dtype not in (torch.float16, torch.bfloat16):
|
if x.dtype not in (torch.float16, torch.bfloat16, torch.float32):
|
||||||
raise ValueError("NVFP4 linear accepts FP16 or BF16 activations.")
|
raise ValueError("NVFP4 linear accepts FP16, BF16, or FP32 activations.")
|
||||||
|
|
||||||
from comfy_kitchen.tensor import QuantizedTensor
|
from comfy_kitchen.tensor import QuantizedTensor
|
||||||
|
|
||||||
|
|
@ -82,6 +82,8 @@ class Nvfp4Linear(nn.Module):
|
||||||
weight = packed_weight.dequantize().to(flat_x)
|
weight = packed_weight.dequantize().to(flat_x)
|
||||||
output = functional.linear(flat_x, weight, bias)
|
output = functional.linear(flat_x, weight, bias)
|
||||||
return output.reshape(*original_shape, self.out_features)
|
return output.reshape(*original_shape, self.out_features)
|
||||||
|
if x.dtype == torch.float32:
|
||||||
|
raise ValueError("Quantized NVFP4 activation GEMM requires FP16 or BF16 activations.")
|
||||||
packed_x = QuantizedTensor.from_float(flat_x, "TensorCoreNVFP4Layout")
|
packed_x = QuantizedTensor.from_float(flat_x, "TensorCoreNVFP4Layout")
|
||||||
output = functional.linear(packed_x, packed_weight, bias)
|
output = functional.linear(packed_x, packed_weight, bias)
|
||||||
return output[:flat_x.shape[0], :self.out_features].reshape(*original_shape, self.out_features)
|
return output[:flat_x.shape[0], :self.out_features].reshape(*original_shape, self.out_features)
|
||||||
|
|
|
||||||
|
|
@ -133,7 +133,7 @@ class Qwen3VL32BTextEncoder(nn.Module):
|
||||||
"""Mounted-checkpoint Qwen3-VL prompt conditioner returning layer-50 states."""
|
"""Mounted-checkpoint Qwen3-VL prompt conditioner returning layer-50 states."""
|
||||||
config = Qwen3VL32BTextConfig()
|
config = Qwen3VL32BTextConfig()
|
||||||
|
|
||||||
def __init__(self, checkpoint_path: str | Path, *, device: str | torch.device = "cuda", dtype: torch.dtype = torch.bfloat16):
|
def __init__(self, checkpoint_path: str | Path, *, device: str | torch.device = "cuda", dtype: torch.dtype = torch.float32):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.checkpoint_path = str(checkpoint_path)
|
self.checkpoint_path = str(checkpoint_path)
|
||||||
self.device_name = str(device)
|
self.device_name = str(device)
|
||||||
|
|
@ -195,7 +195,7 @@ class Qwen3VL32BTextEncoder(nn.Module):
|
||||||
|
|
||||||
class Qwen3VLPromptConditioner:
|
class Qwen3VLPromptConditioner:
|
||||||
"""Tokenize raw H3 prompt text and produce Qwen layer-50 conditioning."""
|
"""Tokenize raw H3 prompt text and produce Qwen layer-50 conditioning."""
|
||||||
def __init__(self, checkpoint_path: str | Path, tokenizer_dir: str | Path | None = None, *, device: str | torch.device = "cuda", dtype: torch.dtype = torch.bfloat16):
|
def __init__(self, checkpoint_path: str | Path, tokenizer_dir: str | Path | None = None, *, device: str | torch.device = "cuda", dtype: torch.dtype = torch.float32):
|
||||||
from .conditioning import H3PromptTokenizer
|
from .conditioning import H3PromptTokenizer
|
||||||
|
|
||||||
tokenizer_dir = tokenizer_dir or Path(__file__).with_name("qwen25_tokenizer")
|
tokenizer_dir = tokenizer_dir or Path(__file__).with_name("qwen25_tokenizer")
|
||||||
|
|
|
||||||
46
tools/compare_qwen0_quantized_dispatch.py
Normal file
46
tools/compare_qwen0_quantized_dispatch.py
Normal file
|
|
@ -0,0 +1,46 @@
|
||||||
|
"""Compare direct NVFP4 wrapper and effective Comfy BF16 weights."""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as functional
|
||||||
|
|
||||||
|
from h3_blackwell_runtime.qwen3vl_text import Qwen3VL32BTextEncoder
|
||||||
|
|
||||||
|
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument("--capture-dir", required=True)
|
||||||
|
parser.add_argument("--checkpoint", default="/text-encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors")
|
||||||
|
parser.add_argument("--dtype", choices=("bfloat16", "float32"), default="bfloat16")
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
encoder = Qwen3VL32BTextEncoder(args.checkpoint, dtype=getattr(torch, args.dtype)).eval()
|
||||||
|
layer = encoder.layers[0]
|
||||||
|
qkv = torch.load(f"{args.capture_dir}/qwen0_quantized_qkv.pt", map_location="cuda", weights_only=False)
|
||||||
|
mlp = torch.load(f"{args.capture_dir}/qwen0_quantized_mlp.pt", map_location="cuda", weights_only=False)
|
||||||
|
|
||||||
|
for group, capture, modules, inputs in (
|
||||||
|
("qkv", qkv, {"q": layer.q_proj, "k": layer.k_proj, "v": layer.v_proj}, {"q": qkv["input"], "k": qkv["input"], "v": qkv["input"]}),
|
||||||
|
("mlp", mlp, {"gate": layer.gate_proj, "up": layer.up_proj, "down": layer.down_proj}, {"gate": mlp["input"], "up": mlp["input"], "down": mlp["activated"]}),
|
||||||
|
):
|
||||||
|
for name, module in modules.items():
|
||||||
|
expected = capture["modules"][name]
|
||||||
|
packed = module._packed_weight()
|
||||||
|
params = packed._params
|
||||||
|
qdata, scale, block_scale = packed.layout_cls.get_plain_tensors(packed)
|
||||||
|
reference_qdata = expected["qdata"].as_subclass(torch.Tensor).cpu()
|
||||||
|
print(f"{group}.{name}.qdata exact={torch.equal(qdata.cpu(), reference_qdata)} shape={tuple(qdata.shape)}")
|
||||||
|
for field, actual in (("scale", scale), ("block_scale", block_scale)):
|
||||||
|
reference = expected[field].to(actual.device)
|
||||||
|
print(f"{group}.{name}.{field} exact={torch.equal(actual, reference)} shape={tuple(actual.shape)} dtype={actual.dtype}")
|
||||||
|
value = inputs[name].to(encoder.dtype)
|
||||||
|
with torch.inference_mode():
|
||||||
|
dequantized = packed.dequantize()
|
||||||
|
direct = module(value)
|
||||||
|
effective = expected["effective_weight"].to("cuda")
|
||||||
|
comfy_effective = functional.linear(value.to(effective.dtype), effective, None)
|
||||||
|
print(f"{group}.{name}.effective dtype={expected['effective_dtype']} shape={expected['effective_shape']} direct_dtype={encoder.dtype}")
|
||||||
|
for variant, actual in (("weight", dequantized), ("linear", direct), ("effective_linear", comfy_effective)):
|
||||||
|
reference = expected["effective_weight"].to(actual.device) if variant == "weight" else expected["output"].to(actual.device)
|
||||||
|
delta = (actual.float() - reference.float()).abs()
|
||||||
|
print(f"{group}.{name}.{variant}_vs_comfy max_abs={delta.max().item():.6g} mean_abs={delta.mean().item():.6g}")
|
||||||
58
tools/patch_comfy_qwen_quantized_tensor_probe.py
Normal file
58
tools/patch_comfy_qwen_quantized_tensor_probe.py
Normal file
|
|
@ -0,0 +1,58 @@
|
||||||
|
"""Capture loaded Comfy Qwen QuantizedTensor fields and dispatch variants."""
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
path = Path("/opt/ComfyUI/comfy/text_encoders/llama.py")
|
||||||
|
source = path.read_text(encoding="utf-8")
|
||||||
|
helper = (
|
||||||
|
"def _h3_owned_cpu(tensor):\n"
|
||||||
|
" return tensor.detach().to(device=\"cpu\", copy=True).contiguous().clone()\n\n"
|
||||||
|
)
|
||||||
|
if "def _h3_owned_cpu(tensor):\n" not in source:
|
||||||
|
source = helper + source
|
||||||
|
old = " xq = self.q_proj(hidden_states)\n xk = self.k_proj(hidden_states)\n xv = self.v_proj(hidden_states)\n"
|
||||||
|
new = old + (
|
||||||
|
" capture_dir = os.getenv(\"H3_CAPTURE_DIR\") if getattr(self, \"_h3_trace_index\", -1) == 0 else None\n"
|
||||||
|
" if capture_dir:\n"
|
||||||
|
" modules = {\"q\": self.q_proj, \"k\": self.k_proj, \"v\": self.v_proj}\n"
|
||||||
|
" def probe(module):\n"
|
||||||
|
" stored = module.weight\n"
|
||||||
|
" qdata, scale, block_scale = stored.layout_cls.get_plain_tensors(stored)\n"
|
||||||
|
" weight, bias, offload_stream = comfy.ops.cast_bias_weight(module, hidden_states, offloadable=True)\n"
|
||||||
|
" try:\n"
|
||||||
|
" output = torch.nn.functional.linear(hidden_states, weight, bias)\n"
|
||||||
|
" return {\"qdata\": _h3_owned_cpu(qdata), \"scale\": _h3_owned_cpu(scale), \"block_scale\": _h3_owned_cpu(block_scale), \"layout_class\": stored._layout_cls, \"orig_dtype\": stored._params.orig_dtype, \"orig_shape\": tuple(stored._params.orig_shape), \"effective_weight\": _h3_owned_cpu(weight), \"effective_dtype\": str(weight.dtype), \"effective_shape\": tuple(weight.shape), \"output\": _h3_owned_cpu(output)}\n"
|
||||||
|
" finally:\n"
|
||||||
|
" comfy.ops.uncast_bias_weight(module, weight, bias, offload_stream)\n"
|
||||||
|
" torch.save({\"input\": hidden_states.detach().cpu(), \"modules\": {name: probe(module) for name, module in modules.items()}}, os.path.join(capture_dir, \"qwen0_quantized_qkv.pt\"))\n"
|
||||||
|
)
|
||||||
|
if source.count(old) != 1:
|
||||||
|
raise RuntimeError("Unable to locate Qwen QKV projections.")
|
||||||
|
source = source.replace(old, new)
|
||||||
|
|
||||||
|
old = " return self.down_proj(self.activation(self.gate_proj(x)) * self.up_proj(x))\n"
|
||||||
|
new = (
|
||||||
|
" gate = self.gate_proj(x)\n"
|
||||||
|
" up = self.up_proj(x)\n"
|
||||||
|
" activated = self.activation(gate) * up\n"
|
||||||
|
" output = self.down_proj(activated)\n"
|
||||||
|
" capture_dir = os.getenv(\"H3_CAPTURE_DIR\") if getattr(self, \"_h3_trace_index\", -1) == 0 else None\n"
|
||||||
|
" if capture_dir:\n"
|
||||||
|
" modules = {\"gate\": self.gate_proj, \"up\": self.up_proj, \"down\": self.down_proj}\n"
|
||||||
|
" def probe(module, value):\n"
|
||||||
|
" stored = module.weight\n"
|
||||||
|
" qdata, scale, block_scale = stored.layout_cls.get_plain_tensors(stored)\n"
|
||||||
|
" weight, bias, offload_stream = comfy.ops.cast_bias_weight(module, value, offloadable=True)\n"
|
||||||
|
" try:\n"
|
||||||
|
" output = torch.nn.functional.linear(value, weight, bias)\n"
|
||||||
|
" return {\"qdata\": _h3_owned_cpu(qdata), \"scale\": _h3_owned_cpu(scale), \"block_scale\": _h3_owned_cpu(block_scale), \"layout_class\": stored._layout_cls, \"orig_dtype\": stored._params.orig_dtype, \"orig_shape\": tuple(stored._params.orig_shape), \"effective_weight\": _h3_owned_cpu(weight), \"effective_dtype\": str(weight.dtype), \"effective_shape\": tuple(weight.shape), \"output\": _h3_owned_cpu(output)}\n"
|
||||||
|
" finally:\n"
|
||||||
|
" comfy.ops.uncast_bias_weight(module, weight, bias, offload_stream)\n"
|
||||||
|
" torch.save({\"input\": x.detach().cpu(), \"activated\": activated.detach().cpu(), \"modules\": {\"gate\": probe(self.gate_proj, x), \"up\": probe(self.up_proj, x), \"down\": probe(self.down_proj, activated)}}, os.path.join(capture_dir, \"qwen0_quantized_mlp.pt\"))\n"
|
||||||
|
" return output\n"
|
||||||
|
)
|
||||||
|
if source.count(old) != 1:
|
||||||
|
raise RuntimeError("Unable to locate Qwen MLP projections.")
|
||||||
|
path.write_text(source.replace(old, new), encoding="utf-8")
|
||||||
|
print("Applied Qwen QuantizedTensor dispatch probe.")
|
||||||
Loading…
Add table
Reference in a new issue