hanjian.thu123 commited on
Commit
5e83b2a
·
1 Parent(s): 64b8c85

[update] fallback slow_attn

Browse files
Files changed (2) hide show
  1. grn/models/basic.py +2 -16
  2. requirements.txt +0 -1
grn/models/basic.py CHANGED
@@ -180,22 +180,8 @@ class SelfAttention(nn.Module):
180
  )
181
  attn_output = attn_output[0].reshape(B, L, C).contiguous()
182
  else:
183
- from kernels import get_kernel
184
- _k = get_kernel("kernels-community/vllm-flash-attn3", revision="main")
185
- flash_attn_varlen_func = _k.flash_attn_varlen_func
186
- attn_output = flash_attn_varlen_func(
187
- q = query_states.squeeze(0),
188
- k = key_states.squeeze(0),
189
- v = value_states.squeeze(0),
190
- cu_seqlens_q=cu_seqlens,
191
- cu_seqlens_k=cu_seqlens,
192
- max_seqlen_q=max_seqlen,
193
- max_seqlen_k=max_seqlen,
194
- softmax_scale=self.scale,
195
- )
196
- attn_output = attn_output.reshape(B, L, C).contiguous()
197
- # # slow attn
198
- # attn_output = slow_attn(query=query_states.transpose(1, 2), key=key_states.transpose(1, 2), value=value_states.transpose(1, 2), scale=self.scale, attn_mask=attn_bias_or_two_vector, dropout_p=0).transpose(1, 2).reshape(B, L, C)
199
 
200
  if sp_manager.sp_on():
201
  # [B, raw_L, C/sp] --> [B, raw_L/sp, C]
 
180
  )
181
  attn_output = attn_output[0].reshape(B, L, C).contiguous()
182
  else:
183
+ # slow attn
184
+ attn_output = slow_attn(query=query_states.transpose(1, 2), key=key_states.transpose(1, 2), value=value_states.transpose(1, 2), scale=self.scale, attn_mask=attn_bias_or_two_vector, dropout_p=0).transpose(1, 2).reshape(B, L, C)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
185
 
186
  if sp_manager.sp_on():
187
  # [B, raw_L, C/sp] --> [B, raw_L/sp, C]
requirements.txt CHANGED
@@ -15,4 +15,3 @@ ftfy>=6.1.1
15
  transformers>=4.35.0
16
  regex>=2023.10.3
17
  pyyaml>=6.0
18
- kernels
 
15
  transformers>=4.35.0
16
  regex>=2023.10.3
17
  pyyaml>=6.0