allow torch.compiling more stuff

This commit is contained in:
kijai
2024-11-13 23:59:10 +02:00
parent acb19bcd29
commit 723537bb75
4 changed files with 26 additions and 14 deletions
@@ -22,7 +22,12 @@ try:
from flash_attn.flash_attn_interface import flash_attn_varlen_func
except:
flash_attn_varlen_func = None
@torch.compiler.disable()
def compute_attention(query, key, value, attn_mask, dropout_p=0.0, is_causal=False):
return F.scaled_dot_product_attention(
query, key, value, dropout_p=dropout_p, is_causal=is_causal, attn_mask=attn_mask,
)
def apply_rope(xq, xk, freqs_cis):
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
@@ -204,7 +209,7 @@ class VarlenSelfAttentionWithT5Mask:
value = value.transpose(1, 2)
# with torch.backends.cuda.sdp_kernel(enable_math=False, enable_flash=False, enable_mem_efficient=True):
stage_hidden_states = F.scaled_dot_product_attention(
stage_hidden_states = compute_attention(
query, key, value, dropout_p=0.0, is_causal=False, attn_mask=attention_mask[i_p],
)
stage_hidden_states = stage_hidden_states.transpose(1, 2).flatten(2, 3) # [bs, tot_seq, dim]
@@ -219,6 +224,7 @@ class VarlenSelfAttentionWithT5Mask:
return output_hidden, output_encoder_hidden
class VarlenFlashSelfAttnSingle:
def __init__(self):
@@ -312,7 +318,7 @@ class VarlenSelfAttnSingle:
key = key.transpose(1, 2).contiguous()
value = value.transpose(1, 2).contiguous()
stage_hidden_states = F.scaled_dot_product_attention(
stage_hidden_states = compute_attention(
query, key, value, dropout_p=0.0, is_causal=False, attn_mask=attention_mask[i_p],
)
stage_hidden_states = stage_hidden_states.transpose(1, 2).flatten(2, 3) # [bs, tot_seq, dim]
@@ -150,7 +150,6 @@ class PyramidDiTForVideoGeneration:
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
return latents
def sample_block_noise(self, bs, ch, temp, height, width):
block_number = bs * ch * temp * (height // 2) * (width // 2)
noise = torch.stack([self.dist.sample() for _ in range(block_number)]) # [block number, 4]