33 lines
1.6 KiB
Python
33 lines
1.6 KiB
Python
|
|
"""Offline Qwen layer-0 attention trace against an existing Comfy capture."""
|
||
|
|
|
||
|
|
import torch
|
||
|
|
import torch.nn.functional as F
|
||
|
|
|
||
|
|
from h3_blackwell_runtime.attention import run_attention
|
||
|
|
from h3_blackwell_runtime.conditioning import H3PromptTokenizer
|
||
|
|
from h3_blackwell_runtime.qwen3vl_text import Qwen3VL32BTextEncoder, _rope
|
||
|
|
|
||
|
|
|
||
|
|
prompt = "A brass-and-paper dragon flies above a rain-washed old city at blue hour."
|
||
|
|
encoder = Qwen3VL32BTextEncoder("/text-encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors", attention_backend="sage2")
|
||
|
|
ids = H3PromptTokenizer("/opt/h3-blackwell-runtime/src/h3_blackwell_runtime/qwen25_tokenizer")(prompt)
|
||
|
|
x = (F.embedding(ids, encoder.embed_tokens).float() * F.embedding(ids, encoder.embed_scale)).to(encoder.dtype)
|
||
|
|
layer = encoder.layers[0]
|
||
|
|
|
||
|
|
with torch.inference_mode():
|
||
|
|
norm = layer.input_layernorm(x)
|
||
|
|
query = layer.q_proj(norm).view(1, 17, 64, 128).transpose(1, 2)
|
||
|
|
key = layer.k_proj(norm).view(1, 17, 8, 128).transpose(1, 2)
|
||
|
|
value = layer.v_proj(norm).view(1, 17, 8, 128).transpose(1, 2)
|
||
|
|
query = layer.q_norm(query)
|
||
|
|
key = layer.k_norm(key)
|
||
|
|
query, key = _rope(query, key, layer.config.rope_theta)
|
||
|
|
key = key.repeat_interleave(8, dim=1)
|
||
|
|
value = value.repeat_interleave(8, dim=1)
|
||
|
|
attention = layer.o_proj(run_attention(query, key, value, backend="sage2", is_causal=True).transpose(1, 2).reshape(1, 17, -1))
|
||
|
|
|
||
|
|
for name, actual in (("norm1", norm), ("attention", attention)):
|
||
|
|
expected = torch.load(f"/capture/qwen0_{name}.pt", map_location="cuda", weights_only=False)
|
||
|
|
delta = (actual.float() - expected.float()).abs()
|
||
|
|
print(f"{name} mean_abs={delta.mean().item():.6g} max_abs={delta.max().item():.6g}")
|