"""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