Capture Qwen projection parity
This commit is contained in:
parent
eaca0034e4
commit
d657dd350a
4 changed files with 84 additions and 9 deletions
19
PARITY.md
19
PARITY.md
|
|
@ -43,6 +43,7 @@ diagnose preview output.
|
||||||
| Qwen attention backend | Comfy source inspection | Qwen chooses Comfy small-input attention, i.e. SDPA path, not Sage2 | Direct Qwen was changed to BF16 causal-mask SDPA with GQA |
|
| Qwen attention backend | Comfy source inspection | Qwen chooses Comfy small-input attention, i.e. SDPA path, not Sage2 | Direct Qwen was changed to BF16 causal-mask SDPA with GQA |
|
||||||
| Qwen NVFP4 metadata | Qwen layer-0 checkpoint sidecars inspected | Every Q/K/V/O and MLP projection has `full_precision_matrix_mult: true`; `o_proj` and `down_proj` also have `pre_quant_scale` | Direct was incorrectly quantizing activations for all Qwen projections |
|
| Qwen NVFP4 metadata | Qwen layer-0 checkpoint sidecars inspected | Every Q/K/V/O and MLP projection has `full_precision_matrix_mult: true`; `o_proj` and `down_proj` also have `pre_quant_scale` | Direct was incorrectly quantizing activations for all Qwen projections |
|
||||||
| Qwen NVFP4 fix | Commit `0cad1db` | Direct honors `full_precision_matrix_mult` and loads/applies `pre_quant_scale`; Spark smoke passed finite expected-shape Q and O projections | Structural fix passed; post-fix numeric QKV/MLP/layer-50 comparisons are still required |
|
| Qwen NVFP4 fix | Commit `0cad1db` | Direct honors `full_precision_matrix_mult` and loads/applies `pre_quant_scale`; Spark smoke passed finite expected-shape Q and O projections | Structural fix passed; post-fix numeric QKV/MLP/layer-50 comparisons are still required |
|
||||||
|
| Qwen layer-0 projections, after NVFP4 fix | Fresh Comfy projection capture at `/tmp/fl2va-qwen0-projections`; direct compared from captured input embedding and post-attention state | Q mean/max `6.6547e-05` / `0.00195485`; K `9.11704e-05` / `0.00241077`; V `6.0298e-05` / `0.00102112`; gate `0.000365695` / `0.0082469`; up `0.000335494` / `0.00792789`; activated `4.76882e-05` / `0.031621`; down `0.000557594` / `0.0749016` | The full-precision metadata fix removed the previous large projection failure. Projection parity is close but not strict; `down_proj` is the largest remaining layer-0 projection boundary. Preserve these numbers and do not recapture this case. |
|
||||||
| H3 DiT block trace | `/tmp/fl2va-full-capture` | Block 0 was close, around mean absolute delta `0.00018`; small independent differences accumulated across blocks | Full 50-block strict parity remains failed/unresolved |
|
| H3 DiT block trace | `/tmp/fl2va-full-capture` | Block 0 was close, around mean absolute delta `0.00018`; small independent differences accumulated across blocks | Full 50-block strict parity remains failed/unresolved |
|
||||||
| Sampler initial state | Comfy `initial.pt` inspected against direct preview | Comfy carries joint AV state; direct preview initializes video noise only and leaves audio zero | Known direct mismatch; identical integer seed is not sampler parity |
|
| Sampler initial state | Comfy `initial.pt` inspected against direct preview | Comfy carries joint AV state; direct preview initializes video noise only and leaves audio zero | Known direct mismatch; identical integer seed is not sampler parity |
|
||||||
| Preview output, before Qwen fixes | Dragon/locomotive prompt previews | Flower-like output | Confirmed incorrect conditioning had visible semantic impact |
|
| Preview output, before Qwen fixes | Dragon/locomotive prompt previews | Flower-like output | Confirmed incorrect conditioning had visible semantic impact |
|
||||||
|
|
@ -67,13 +68,14 @@ diagnose preview output.
|
||||||
Only these are outstanding. Do not recapture or revisit rows marked complete
|
Only these are outstanding. Do not recapture or revisit rows marked complete
|
||||||
unless the checkpoint, Comfy version, prompt, or backend changes.
|
unless the checkpoint, Comfy version, prompt, or backend changes.
|
||||||
|
|
||||||
1. **Qwen projection replay after `0cad1db`.** Use the existing immutable
|
1. **Qwen all-layer replay after `0cad1db`.** The layer-0 projection capture
|
||||||
layer-0 capture. Compare direct Q/K/V, O, gate, up, activated, and down
|
was completed: errors are recorded above. Replay the existing 50 layer
|
||||||
against Comfy. This is the first unresolved boundary after the completed
|
tensors in `/tmp/fl2va-qwen-trace`; report first failing layer and mean/max
|
||||||
embedding/RMSNorm/backend investigations.
|
error. Do not generate another video before this is recorded.
|
||||||
2. **Qwen all-layer replay after the projection gate passes.** Replay the
|
2. **Qwen projection precision follow-up if layer replay fails at layer 0.**
|
||||||
existing 50 layer tensors in `/tmp/fl2va-qwen-trace`; report first failing
|
Do not repeat the capture. Use `/tmp/fl2va-qwen0-projections` to isolate the
|
||||||
layer and mean/max error. Do not generate another video before this passes.
|
remaining `down_proj` max error `0.0749016` and its BF16/dequantization
|
||||||
|
boundary.
|
||||||
3. **Token-refiner capture and comparison.** Capture exact Comfy pre-refiner,
|
3. **Token-refiner capture and comparison.** Capture exact Comfy pre-refiner,
|
||||||
block 0, block 1, and final refined states once, then compare direct using
|
block 0, block 1, and final refined states once, then compare direct using
|
||||||
the matched Qwen layer-50 tensor.
|
the matched Qwen layer-50 tensor.
|
||||||
|
|
@ -101,7 +103,8 @@ unless the checkpoint, Comfy version, prompt, or backend changes.
|
||||||
|
|
||||||
| Gate | Tools |
|
| Gate | Tools |
|
||||||
| --- | --- |
|
| --- | --- |
|
||||||
| Qwen captures | `patch_comfy_qwen_output_capture.py`, `patch_comfy_qwen_layer_trace.py`, `patch_comfy_qwen_layer0_sublayers.py` |
|
| Qwen captures | `patch_comfy_qwen_output_capture.py`, `patch_comfy_qwen_layer_trace.py`, `patch_comfy_qwen_layer0_sublayers.py`, `patch_comfy_qwen_layer0_projections.py` |
|
||||||
|
| Qwen projection gate | `compare_qwen0_projections.py` |
|
||||||
| Qwen local diagnosis | `trace_qwen0_attention.py` needs repair before reuse; it currently passes an obsolete encoder argument |
|
| Qwen local diagnosis | `trace_qwen0_attention.py` needs repair before reuse; it currently passes an obsolete encoder argument |
|
||||||
| H3 block capture | `patch_comfy_h3_capture.py`, `patch_comfy_h3_block_capture.py`, `patch_comfy_h3_block0_sublayers.py` |
|
| H3 block capture | `patch_comfy_h3_capture.py`, `patch_comfy_h3_block_capture.py`, `patch_comfy_h3_block0_sublayers.py` |
|
||||||
| H3 local diagnosis | `compare_block0_qkv.py`, `compare_attention_backends.py`, `compare_block0_mlp_projections.py`, `localize_block_sublayers.py`, `localize_block_mismatch.py` |
|
| H3 local diagnosis | `compare_block0_qkv.py`, `compare_attention_backends.py`, `compare_block0_mlp_projections.py`, `localize_block_sublayers.py`, `localize_block_mismatch.py` |
|
||||||
|
|
|
||||||
35
tools/compare_qwen0_projections.py
Normal file
35
tools/compare_qwen0_projections.py
Normal file
|
|
@ -0,0 +1,35 @@
|
||||||
|
"""Compare direct Qwen layer-0 projections with a Comfy capture."""
|
||||||
|
|
||||||
|
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")
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
encoder = Qwen3VL32BTextEncoder(args.checkpoint).eval()
|
||||||
|
layer = encoder.layers[0]
|
||||||
|
embeds = torch.load(f"{args.capture_dir}/qwen_input_embeds.pt", map_location="cuda", weights_only=False)
|
||||||
|
qkv_expected = torch.load(f"{args.capture_dir}/qwen0_qkv.pt", map_location="cuda", weights_only=False)
|
||||||
|
mlp_expected = torch.load(f"{args.capture_dir}/qwen0_mlp_projections.pt", map_location="cuda", weights_only=False)
|
||||||
|
|
||||||
|
with torch.inference_mode():
|
||||||
|
norm1 = layer.input_layernorm(embeds.to(encoder.dtype))
|
||||||
|
qkv_actual = {"q": layer.q_proj(norm1), "k": layer.k_proj(norm1), "v": layer.v_proj(norm1)}
|
||||||
|
post_attention = torch.load(f"{args.capture_dir}/qwen0_post_attention.pt", map_location="cuda", weights_only=False)
|
||||||
|
norm2 = layer.post_attention_layernorm(post_attention.to(encoder.dtype))
|
||||||
|
gate = layer.gate_proj(norm2)
|
||||||
|
up = layer.up_proj(norm2)
|
||||||
|
activated = functional.silu(gate) * up
|
||||||
|
mlp_actual = {"gate": gate, "up": up, "activated": activated, "down": layer.down_proj(activated)}
|
||||||
|
|
||||||
|
for group, actual, expected in (("qkv", qkv_actual, qkv_expected), ("mlp", mlp_actual, mlp_expected)):
|
||||||
|
for name, value in actual.items():
|
||||||
|
delta = (value.float() - expected[name].float()).abs()
|
||||||
|
print(f"{group}.{name} shape={tuple(value.shape)} max_abs={delta.max().item():.6g} mean_abs={delta.mean().item():.6g}")
|
||||||
32
tools/patch_comfy_qwen_layer0_projections.py
Normal file
32
tools/patch_comfy_qwen_layer0_projections.py
Normal file
|
|
@ -0,0 +1,32 @@
|
||||||
|
"""Capture Comfy Qwen layer-0 projection boundaries for direct parity checks."""
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
path = Path("/opt/ComfyUI/comfy/text_encoders/llama.py")
|
||||||
|
source = path.read_text(encoding="utf-8")
|
||||||
|
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"
|
||||||
|
" torch.save({\"q\": xq.detach().cpu(), \"k\": xk.detach().cpu(), \"v\": xv.detach().cpu()}, os.path.join(capture_dir, \"qwen0_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"
|
||||||
|
" torch.save({\"gate\": gate.detach().cpu(), \"up\": up.detach().cpu(), \"activated\": activated.detach().cpu(), \"down\": output.detach().cpu()}, os.path.join(capture_dir, \"qwen0_mlp_projections.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 layer-0 projection capture patch.")
|
||||||
|
|
@ -57,7 +57,12 @@ if source.count(old) != 1:
|
||||||
source = source.replace(old, new)
|
source = source.replace(old, new)
|
||||||
|
|
||||||
old = " for i, layer in enumerate(self.layers):\n"
|
old = " for i, layer in enumerate(self.layers):\n"
|
||||||
new = " for i, layer in enumerate(self.layers):\n layer._h3_trace_index = i\n"
|
new = (
|
||||||
|
" for i, layer in enumerate(self.layers):\n"
|
||||||
|
" layer._h3_trace_index = i\n"
|
||||||
|
" layer.self_attn._h3_trace_index = i\n"
|
||||||
|
" layer.mlp._h3_trace_index = i\n"
|
||||||
|
)
|
||||||
if source.count(old) != 1:
|
if source.count(old) != 1:
|
||||||
raise RuntimeError("Unable to locate Qwen decoder layer loop.")
|
raise RuntimeError("Unable to locate Qwen decoder layer loop.")
|
||||||
path.write_text(source.replace(old, new), encoding="utf-8")
|
path.write_text(source.replace(old, new), encoding="utf-8")
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue