ProCreations's picture
Accelerate full 40-step FP8 generation with native precision, measured quality and real-time demo
1081be0 verified
Raw History Blame Contribute Delete
4.13 kB
"""Full-compute acceleration for Image 2.1 Calibrated FP8. Built with Qwen.
Keeps every denoising step, BF16 attention, original calibrated FP8 GEMMs and
FP32 accumulation/scales. No approximate residual cache or attention quantization.
Only cached-prefix decode blocks compile; prefill keeps upstream behavior.
"""
import types
import torch
from diffusers.models.transformers.transformer_qwenimage21 import (
QwenImage21AttnProcessor, QwenImage21TransformerBlock,
)
_ORIGINAL_BLOCK_FORWARD = QwenImage21TransformerBlock.forward
def _real_rope(x, frequencies):
paired = x.float().unflatten(-1, (-1, 2))
cosine = frequencies.real[None, :, None, :]
sine = frequencies.imag[None, :, None, :]
return torch.stack((
paired[..., 0] * cosine - paired[..., 1] * sine,
paired[..., 0] * sine + paired[..., 1] * cosine,
), dim=-1).flatten(-2).to(x.dtype)
class NativeAttentionProcessor(QwenImage21AttnProcessor):
def __call__(self, attn, hidden_states, attention_mask=None, rotary_emb=None,
layer_cache=None, kv_cache_mode=None, cache_write_slice=None,
segments=None, key_valid=None):
if kv_cache_mode != 'cached' or attention_mask is not None:
return super().__call__(attn, hidden_states, attention_mask, rotary_emb,
layer_cache, kv_cache_mode, cache_write_slice,
segments, key_valid)
query = attn.to_q(hidden_states).unflatten(-1, (attn.heads, -1))
key = attn.to_k(hidden_states).unflatten(-1, (attn.heads, -1))
value = attn.to_v(hidden_states).unflatten(-1, (attn.heads, -1))
query = attn.norm_q(query)
key = attn.norm_k(key)
if rotary_emb is not None:
query = _real_rope(query, rotary_emb)
key = _real_rope(key, rotary_emb)
cached_key, cached_value = layer_cache.get()
key = torch.cat((cached_key, key), dim=1)
value = torch.cat((cached_value, value), dim=1)
output = torch.nn.functional.scaled_dot_product_attention(
query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2)
).transpose(1, 2)
return attn.to_out[1](attn.to_out[0](output.flatten(2, 3)))
def _decode_block(self, hidden_states, modulation, rotary_emb=None,
attention_mask=None, target_token_mask=None, layer_cache=None,
kv_cache_mode=None, cache_write_slice=None, segments=None,
key_valid=None):
# Cached decode contains only target image tokens. The final t=0 modulation
# row belongs to the prefix, which was already evaluated during prefill.
if kv_cache_mode == 'cached':
modulation = modulation[:-1]
target_token_mask = None
return _ORIGINAL_BLOCK_FORWARD(
self, hidden_states, modulation, rotary_emb, attention_mask,
target_token_mask, layer_cache, kv_cache_mode, cache_write_slice,
segments, key_valid,
)
def _dispatch_block(self, **kwargs):
if kwargs.get('kv_cache_mode') == 'cached':
return self._image21_compiled(**kwargs)
return _ORIGINAL_BLOCK_FORWARD(self, **kwargs)
def accelerate_pipeline(pipe):
"""Enable once after loading. Initial compilation is excluded from warm timings.
Dynamic sequence lengths reduce recompilation across prompts/resolutions.
emulate_precision_casts preserves the upstream intermediate BF16 rounding
boundaries in fused code; GPU reduction ordering can still differ.
"""
if getattr(pipe, '_image21_accelerated', False):
return pipe
torch._dynamo.config.recompile_limit = max(torch._dynamo.config.recompile_limit, 64)
for block in pipe.transformer.transformer_blocks:
block.attn.set_processor(NativeAttentionProcessor())
block._image21_compiled = torch.compile(
types.MethodType(_decode_block, block), fullgraph=True, dynamic=True,
options={'emulate_precision_casts': True},
)
block.forward = types.MethodType(_dispatch_block, block)
pipe._image21_accelerated = True
return pipe