File size: 4,131 Bytes
1081be0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
"""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