XAFT commited on
Commit
102efd5
·
verified ·
1 Parent(s): ff7a5d8

Add support for FlashAttention

Browse files
Files changed (3) hide show
  1. README.md +3 -0
  2. modeling_selectivevit.py +2 -0
  3. selective_vit.py +173 -51
README.md CHANGED
@@ -54,6 +54,7 @@ This model is intended for:
54
  ### Image Classification Example
55
 
56
  ```python
 
57
  from transformers import AutoModelForImageClassification, AutoImageProcessor
58
  from PIL import Image
59
  import requests
@@ -72,12 +73,14 @@ model = AutoModelForImageClassification.from_pretrained(
72
  "XAFT/SM-Selective-ViT-Tiny-Tall-224",
73
  trust_remote_code=True,
74
  )
 
75
 
76
  # Preprocess
77
  inputs = processor(
78
  images=image,
79
  return_tensors="pt",
80
  )
 
81
 
82
  # Forward pass
83
  outputs = model(**inputs)
 
54
  ### Image Classification Example
55
 
56
  ```python
57
+ import torch
58
  from transformers import AutoModelForImageClassification, AutoImageProcessor
59
  from PIL import Image
60
  import requests
 
73
  "XAFT/SM-Selective-ViT-Tiny-Tall-224",
74
  trust_remote_code=True,
75
  )
76
+ model = model.half() # Cast to FP16 to enable FlashAttention
77
 
78
  # Preprocess
79
  inputs = processor(
80
  images=image,
81
  return_tensors="pt",
82
  )
83
+ inputs = inputs.to(torch.half) # Cast to FP16
84
 
85
  # Forward pass
86
  outputs = model(**inputs)
modeling_selectivevit.py CHANGED
@@ -50,6 +50,7 @@ class SMSelectiveViTModelForClassification(PreTrainedModel ):
50
  full=False,
51
  output_hidden_states=None,
52
  return_dict=None,
 
53
  **kwargs,
54
  ):
55
  output_hidden_states = (
@@ -67,6 +68,7 @@ class SMSelectiveViTModelForClassification(PreTrainedModel ):
67
  pixel_values,
68
  full=full,
69
  output_hidden_states=output_hidden_states,
 
70
  )
71
 
72
  logits, distil_logits = self.backbone.forward_classifier(last_hidden)
 
50
  full=False,
51
  output_hidden_states=None,
52
  return_dict=None,
53
+ skip_masks=False,
54
  **kwargs,
55
  ):
56
  output_hidden_states = (
 
68
  pixel_values,
69
  full=full,
70
  output_hidden_states=output_hidden_states,
71
+ skip_masks=skip_masks
72
  )
73
 
74
  logits, distil_logits = self.backbone.forward_classifier(last_hidden)
selective_vit.py CHANGED
@@ -1,8 +1,23 @@
 
1
  import torch
2
  import torch.nn as nn
3
  import torch.nn.functional as F
4
  from torchvision.ops import StochasticDepth
5
  import math
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6
 
7
  class SoftMaskedMultiheadAttention(nn.Module):
8
  def __init__(self, embed_dim, num_heads, dropout=0.0, bias=True,
@@ -47,12 +62,9 @@ class SoftMaskedMultiheadAttention(nn.Module):
47
  if self.v_proj.bias is not None:
48
  nn.init.constant_(self.out_proj.bias, 0.)
49
 
50
- def forward(self, query, key, value, key_padding_mask=None,
51
- need_weights=True, attn_mask=None, average_attn_weights=True):
52
- """
53
- query, key, value: shape (L, N, E)
54
- where L is the sequence length, N is the batch size, E is the embedding dimension.
55
- """
56
  batch_size, tgt_len, embed_dim = query.size()
57
  batch_size, src_len, _ = key.size()
58
 
@@ -74,10 +86,11 @@ class SoftMaskedMultiheadAttention(nn.Module):
74
  # Ensure attn_mask values are in (0, 1] to avoid log(0)
75
  # attn_mask shape [b, l]
76
  attn_mask = attn_mask.unsqueeze(1).unsqueeze(1)
77
- if not self.training:
78
- scores = scores.masked_fill((attn_mask == 0.), float('-inf'))
79
  eps = 1e-6
80
- attn_mask = attn_mask.clip(min=eps).log()
 
 
 
81
  # attn_mask shape [b, 1, 1, l]
82
  scores = scores + self.scale * attn_mask
83
 
@@ -99,16 +112,96 @@ class SoftMaskedMultiheadAttention(nn.Module):
99
 
100
  attn_output = self.out_proj(attn_output)
101
 
102
- if need_weights:
103
- # Optionally average attention weights over heads
104
- if average_attn_weights:
105
- attn_weights = attn_weights.mean(dim=1)
106
- else:
107
- attn_weights = attn_weights
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
108
  else:
109
- attn_weights = None
 
110
 
111
- return attn_output, attn_weights
112
 
113
  def get_ffn(input_dim, output_dim, middle_dim, dropout=0.1):
114
  fc1 = nn.Linear(input_dim, middle_dim)
@@ -122,6 +215,7 @@ def get_ffn(input_dim, output_dim, middle_dim, dropout=0.1):
122
  nn.Dropout(dropout),
123
  fc3
124
  )
 
125
  # Assuming SoftMaskedMultiheadAttention is already defined as provided earlier
126
  class EncoderBlock(nn.Module):
127
  def __init__(self, input_dim, embed_dim, num_heads, mlp_dim, dropout=0.1, drop_path=0.0, patch_drop=0.0, attention_scale=2., mask_threshold=0.05):
@@ -162,16 +256,13 @@ class EncoderBlock(nn.Module):
162
  nn.init.zeros_(self.norm3.weight)
163
  nn.init.zeros_(self.norm3.bias)
164
 
165
- def forward_common(self, x, mask):
166
- """
167
- x: shape (batch_size, seq_len, embed_dim)
168
- """
169
  # Compute mask scores: (batch_size, seq_len, 1)
170
  x1 = x
171
  x = self.embed(x)
172
  x = self.norm1(x)
173
  # Apply attention mechanism
174
- attn_output, _ = self.self_attn(x, x, x, attn_mask=mask)
175
  # Add & Norm
176
  x = x + self.path_drop(attn_output)
177
  x = self.norm2(x)
@@ -185,6 +276,45 @@ class EncoderBlock(nn.Module):
185
  x = x1 + x
186
  return x
187
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
188
  def get_groups(self, mask, full=False):
189
  n_items, index = (mask != 0.0).sum(-1).cpu().sort(descending=True)
190
  n_items, index = n_items.tolist(), index.tolist()
@@ -198,13 +328,7 @@ class EncoderBlock(nn.Module):
198
  groups[-1][0].append(ii)
199
  return groups
200
 
201
- def infer_forward(self, x, mask, full=False):
202
- """
203
- The “sparse‐inference” path: for each group of batch‐samples that have the same
204
- number n of tokens ≥ mask_threshold, gather only those top‐n tokens (in original order),
205
- run forward_common on the smaller (b’, n, dim) tensor, then scatter the results back.
206
- Fully masked tokens are left untouched.
207
- """
208
  # Step 1: Threshold the mask without in-place ops
209
  mask_thresholded = mask * (mask >= self.mask_threshold)
210
  # Step 2: Prepare output tensor (copy of x)
@@ -223,7 +347,7 @@ class EncoderBlock(nn.Module):
223
  X_topk = torch.gather(x_sel, dim=1, index=idx_expanded)
224
  mask_topk = torch.gather(mask_sel, dim=1, index=topk_idx_sorted)
225
  # Run forward pass
226
- results = self.forward_common(X_topk, mask_topk)
227
  # Scatter results into a new x_sel tensor
228
  x_sel_updated = x_sel.clone()
229
  x_sel_updated = x_sel_updated.scatter(1, idx_expanded, results)
@@ -231,15 +355,27 @@ class EncoderBlock(nn.Module):
231
  x_out[batch_indices] = x_sel_updated
232
  return x_out
233
 
234
- def forward(self, x, full=False):
235
  if self.linear_mask is not None:
236
  attn_mask = self.patch_drop(self.linear_mask(x).sigmoid().squeeze(-1))
237
  else:
238
  attn_mask = None
239
  if not self.training and not attn_mask is None and self.mask_threshold >= 0:
240
- x = self.infer_forward(x, attn_mask, full)
 
 
 
 
 
 
 
 
 
 
 
 
241
  else:
242
- x = self.forward_common(x, attn_mask)
243
  return x, attn_mask
244
 
245
 
@@ -329,15 +465,8 @@ class VisionTransformer(nn.Module):
329
  pixel_values,
330
  full=False,
331
  output_hidden_states=False,
 
332
  ):
333
- """
334
- Args:
335
- pixel_values: (B, C, H, W)
336
- Returns:
337
- last_hidden_state: (B, N, D)
338
- all_hidden_states: tuple or None
339
- masks: Tensor or None
340
- """
341
  batch_size = pixel_values.size(0)
342
  hidden_states = []
343
 
@@ -361,7 +490,7 @@ class VisionTransformer(nn.Module):
361
  masks = []
362
 
363
  for layer in self.encoder_layers:
364
- x, mask = layer(x, full)
365
 
366
  if output_hidden_states:
367
  hidden_states.append(x)
@@ -385,13 +514,6 @@ class VisionTransformer(nn.Module):
385
 
386
 
387
  def forward_classifier(self, hidden_states):
388
- """
389
- Args:
390
- hidden_states: (B, N, D)
391
- Returns:
392
- logits: (B, num_classes)
393
- dis_logits: (B, num_classes) or None
394
- """
395
  cls_token = hidden_states[:, 0]
396
  logits = self.head(cls_token)
397
 
@@ -406,7 +528,7 @@ class VisionTransformer(nn.Module):
406
 
407
  return logits, dis_logits
408
 
409
- def forward(self, x, full=False):
410
- last_hidden_states, hidden_states, masks = self.forward_features(x, full)
411
  logits, dis_logits = self.forward_classifier(last_hidden_states)
412
  return logits, dis_logits, masks
 
1
+ import os
2
  import torch
3
  import torch.nn as nn
4
  import torch.nn.functional as F
5
  from torchvision.ops import StochasticDepth
6
  import math
7
+ import warnings
8
+ try:
9
+ import torch.nn.attention.varlen as varlen
10
+ HAS_VARLEN_FLASH_ATTENTION = True
11
+ except ImportError:
12
+ warnings.warn(
13
+ "Could not import torch.nn.attention.varlen, variable length Flash Attention is disabled.",
14
+ category=UserWarning,
15
+ stacklevel=2,
16
+ )
17
+ HAS_VARLEN_FLASH_ATTENTION = False
18
+
19
+ enable_fa = os.environ.get('DISABLE_FA', '0').lower() not in {"1", "true", "yes", "y", "on"}
20
+ HAS_VARLEN_FLASH_ATTENTION = HAS_VARLEN_FLASH_ATTENTION and enable_fa
21
 
22
  class SoftMaskedMultiheadAttention(nn.Module):
23
  def __init__(self, embed_dim, num_heads, dropout=0.0, bias=True,
 
62
  if self.v_proj.bias is not None:
63
  nn.init.constant_(self.out_proj.bias, 0.)
64
 
65
+
66
+ def naive_forward(self, query, key, value, key_padding_mask=None,
67
+ attn_mask=None, average_attn_weights=True):
 
 
 
68
  batch_size, tgt_len, embed_dim = query.size()
69
  batch_size, src_len, _ = key.size()
70
 
 
86
  # Ensure attn_mask values are in (0, 1] to avoid log(0)
87
  # attn_mask shape [b, l]
88
  attn_mask = attn_mask.unsqueeze(1).unsqueeze(1)
 
 
89
  eps = 1e-6
90
+ attn_mask_l = attn_mask.clip(min=eps).log()
91
+ if not self.training:
92
+ attn_mask_l = attn_mask_l.masked_fill((attn_mask == 0.), float('-inf'))
93
+ attn_mask = attn_mask_l
94
  # attn_mask shape [b, 1, 1, l]
95
  scores = scores + self.scale * attn_mask
96
 
 
112
 
113
  attn_output = self.out_proj(attn_output)
114
 
115
+ return attn_output
116
+
117
+ def flash_forward(
118
+ self,
119
+ query, key, value,
120
+ cu_seq_q, cu_seq_k,
121
+ max_q, max_k,
122
+ attn_mask=None,
123
+ is_causal=False,
124
+ ):
125
+ """
126
+ FlashAttention-compatible soft-masked attention using varlen_attn
127
+ """
128
+
129
+ q = self.q_proj(query) # (Tq, H*D)
130
+ k = self.k_proj(key) # (Tk, H*D)
131
+ v = self.v_proj(value) # (Tk, H*D)
132
+
133
+ Tq = q.shape[0]
134
+ Tk = k.shape[0]
135
+
136
+ q = q.view(Tq, self.num_heads, self.head_dim)
137
+ k = k.view(Tk, self.num_heads, self.head_dim)
138
+ v = v.view(Tk, self.num_heads, self.head_dim)
139
+
140
+ # Apply the soft [0, 1] mask
141
+ if attn_mask is not None:
142
+ # attn_mask: (Tk,) or (B, Lk) flattened to match Tk
143
+ eps = 1e-6
144
+ attn_mask_l = attn_mask.clip(min=eps).log()
145
+ if not self.training: # Inference mode can have infinite atten scores
146
+ attn_mask_l = attn_mask_l.masked_fill((attn_mask == 0.), float('-inf'))
147
+ log_m = attn_mask_l
148
+
149
+ # Broadcast to (Tk, H, 1)
150
+ log_m = log_m.view(Tk, 1, 1).expand(-1, self.num_heads, 1)
151
+ k_zeros = torch.zeros_like(log_m).expand(-1, -1, 7)
152
+
153
+ # Augment K and Q
154
+ # We want:
155
+ # (qk^T)/sqrt(d) + scale * log(m)
156
+ scale_attn = 1.0 / math.sqrt(self.head_dim)
157
+
158
+ k_extra = log_m * (self.scale / scale_attn)
159
+
160
+ k = torch.cat([k, k_extra, k_zeros], dim=-1) # (Tk, H, D+1)
161
+
162
+ v_zeros = torch.zeros(Tk, self.num_heads, 8, device=v.device, dtype=v.dtype)
163
+ v = torch.cat([v, v_zeros], dim=-1)
164
+
165
+ q_ones = torch.ones(
166
+ Tq, self.num_heads, 8,
167
+ device=q.device, dtype=q.dtype
168
+ )
169
+ q = torch.cat([q, q_ones], dim=-1) # (Tq, H, D+1)
170
+
171
+ attn_dim = self.head_dim + 1
172
+ else:
173
+ attn_dim = self.head_dim
174
+ scale_attn = 1.0 / math.sqrt(self.head_dim)
175
+
176
+ # FlashAttention varlen call
177
+ out = varlen.varlen_attn(
178
+ query=q,
179
+ key=k,
180
+ value=v,
181
+ cu_seq_q=cu_seq_q,
182
+ cu_seq_k=cu_seq_k,
183
+ max_q=max_q,
184
+ max_k=max_k,
185
+ is_causal=is_causal,
186
+ scale=scale_attn,
187
+ )
188
+
189
+ # Merge heads and output projection
190
+ out = out[..., :self.head_dim]
191
+ out = out.reshape(Tq, self.num_heads * self.head_dim)
192
+ out = self.out_proj(out)
193
+
194
+ return out
195
+
196
+ def forward(self, query, key, value, method="naive", **kwargs):
197
+ if method == 'naive':
198
+ out = self.naive_forward(query, key, value, **kwargs)
199
+ elif method == "fa":
200
+ out = self.flash_forward(query, key, value, **kwargs)
201
  else:
202
+ raise ValueError(f"No attention method named {method}.")
203
+ return out
204
 
 
205
 
206
  def get_ffn(input_dim, output_dim, middle_dim, dropout=0.1):
207
  fc1 = nn.Linear(input_dim, middle_dim)
 
215
  nn.Dropout(dropout),
216
  fc3
217
  )
218
+
219
  # Assuming SoftMaskedMultiheadAttention is already defined as provided earlier
220
  class EncoderBlock(nn.Module):
221
  def __init__(self, input_dim, embed_dim, num_heads, mlp_dim, dropout=0.1, drop_path=0.0, patch_drop=0.0, attention_scale=2., mask_threshold=0.05):
 
256
  nn.init.zeros_(self.norm3.weight)
257
  nn.init.zeros_(self.norm3.bias)
258
 
259
+ def forward_common(self, x, mask, skip_masks=False):
 
 
 
260
  # Compute mask scores: (batch_size, seq_len, 1)
261
  x1 = x
262
  x = self.embed(x)
263
  x = self.norm1(x)
264
  # Apply attention mechanism
265
+ attn_output = self.self_attn(x, x, x, attn_mask=mask if not skip_masks else None, method="naive")
266
  # Add & Norm
267
  x = x + self.path_drop(attn_output)
268
  x = self.norm2(x)
 
276
  x = x1 + x
277
  return x
278
 
279
+ def flash_forward(self, x, mask, skip_masks=False):
280
+ binary_mask = mask >= self.mask_threshold
281
+ sel_mask = mask[binary_mask]
282
+
283
+ seq_lengths = binary_mask.sum(1)
284
+ cum_lengths = torch.zeros(binary_mask.shape[0]+1, dtype=torch.int, device=binary_mask.device)
285
+ cum_lengths[1:] = seq_lengths.cumsum(-1)
286
+ max_len = seq_lengths.amax()
287
+
288
+ x1 = x
289
+ x = x[binary_mask]
290
+
291
+ x = self.embed(x)
292
+ x = self.norm1(x)
293
+ # Apply flash attention mechanism
294
+ attn_output = self.self_attn(
295
+ x, x, x,
296
+ cu_seq_q=cum_lengths,
297
+ cu_seq_k=cum_lengths,
298
+ max_q=max_len,
299
+ max_k=max_len,
300
+ attn_mask=sel_mask if not skip_masks else None,
301
+ method="fa"
302
+ )
303
+ # Add & Norm
304
+ x = x + self.path_drop(attn_output)
305
+ x = self.norm2(x)
306
+ # Feed-forward network
307
+ mlp_output = self.mlp(x)
308
+ # Add & Norm
309
+ x = self.path_drop(self.project(x + mlp_output))
310
+ x = self.norm3(x)
311
+ if mask is not None:
312
+ x = x * sel_mask.unsqueeze(-1)
313
+
314
+ x_out = x1.clone()
315
+ x_out[binary_mask] = x_out[binary_mask] + x
316
+ return x_out
317
+
318
  def get_groups(self, mask, full=False):
319
  n_items, index = (mask != 0.0).sum(-1).cpu().sort(descending=True)
320
  n_items, index = n_items.tolist(), index.tolist()
 
328
  groups[-1][0].append(ii)
329
  return groups
330
 
331
+ def naive_forward(self, x, mask, full=False, skip_masks=False):
 
 
 
 
 
 
332
  # Step 1: Threshold the mask without in-place ops
333
  mask_thresholded = mask * (mask >= self.mask_threshold)
334
  # Step 2: Prepare output tensor (copy of x)
 
347
  X_topk = torch.gather(x_sel, dim=1, index=idx_expanded)
348
  mask_topk = torch.gather(mask_sel, dim=1, index=topk_idx_sorted)
349
  # Run forward pass
350
+ results = self.forward_common(X_topk, mask_topk, skip_masks)
351
  # Scatter results into a new x_sel tensor
352
  x_sel_updated = x_sel.clone()
353
  x_sel_updated = x_sel_updated.scatter(1, idx_expanded, results)
 
355
  x_out[batch_indices] = x_sel_updated
356
  return x_out
357
 
358
+ def forward(self, x, full=False, skip_masks=False):
359
  if self.linear_mask is not None:
360
  attn_mask = self.patch_drop(self.linear_mask(x).sigmoid().squeeze(-1))
361
  else:
362
  attn_mask = None
363
  if not self.training and not attn_mask is None and self.mask_threshold >= 0:
364
+ if (
365
+ HAS_VARLEN_FLASH_ATTENTION and
366
+ 'cuda' in x.device.type and
367
+ x.dtype in (torch.bfloat16, torch.float16)
368
+ ):
369
+ x = self.flash_forward(x, attn_mask, skip_masks)
370
+ else:
371
+ warnings.warn(
372
+ "Flash Attention requirements not met, falling back to naive attention.",
373
+ category=UserWarning,
374
+ stacklevel=2,
375
+ )
376
+ x = self.naive_forward(x, attn_mask, full, skip_masks)
377
  else:
378
+ x = self.forward_common(x, attn_mask, skip_masks)
379
  return x, attn_mask
380
 
381
 
 
465
  pixel_values,
466
  full=False,
467
  output_hidden_states=False,
468
+ skip_masks=False
469
  ):
 
 
 
 
 
 
 
 
470
  batch_size = pixel_values.size(0)
471
  hidden_states = []
472
 
 
490
  masks = []
491
 
492
  for layer in self.encoder_layers:
493
+ x, mask = layer(x, full, skip_masks=skip_masks)
494
 
495
  if output_hidden_states:
496
  hidden_states.append(x)
 
514
 
515
 
516
  def forward_classifier(self, hidden_states):
 
 
 
 
 
 
 
517
  cls_token = hidden_states[:, 0]
518
  logits = self.head(cls_token)
519
 
 
528
 
529
  return logits, dis_logits
530
 
531
+ def forward(self, x, full=False, skip_masks=False):
532
+ last_hidden_states, hidden_states, masks = self.forward_features(x, full, skip_masks=skip_masks)
533
  logits, dis_logits = self.forward_classifier(last_hidden_states)
534
  return logits, dis_logits, masks