From 10d46493cb8d2f50eadf59956948afdc15f56a20 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Wed, 6 Nov 2024 15:39:32 -0800 Subject: [PATCH] Update flex_attention.py --- unsloth/kernels/flex_attention.py | 16 +++++----------- 1 file changed, 5 insertions(+), 11 deletions(-) diff --git a/unsloth/kernels/flex_attention.py b/unsloth/kernels/flex_attention.py index dfd48504d5..678574928c 100644 --- a/unsloth/kernels/flex_attention.py +++ b/unsloth/kernels/flex_attention.py @@ -42,14 +42,15 @@ if not HAS_FLEX_ATTENTION: # Below fails on compiled_autograd, so disable it try: - old_compiled_autograd = torch._dynamo.config.compiled_autograd - torch._dynamo.config.compiled_autograd = False + disable_compiled_autograd = torch._dynamo.compiled_autograd.disable except: - old_compiled_autograd = False + disable_compiled_autograd = lambda *args, **kwargs: *args, **kwargs pass # Logit softcapping - @torch.compile(fullgraph = True, dynamic = True, options = torch_compile_options) + @disable_compiled_autograd( + torch.compile(fullgraph = True, dynamic = True, options = torch_compile_options) + ) def slow_attention_softcapping(Q, K, V, causal_mask, self, bsz, q_len): n_heads = self.num_heads head_dim = self.head_dim @@ -82,13 +83,6 @@ if not HAS_FLEX_ATTENTION: return A pass - # Return compiled_autograd back - try: - torch._dynamo.config.compiled_autograd = old_compiled_autograd - except: - pass - pass - create_flex_attention_causal_mask = None create_flex_attention_sliding_window_mask = None else: