ereniko commited on
Commit
3e6f95c
·
verified ·
1 Parent(s): 3fb0189

Delete attention.py

Browse files
Files changed (1) hide show
  1. attention.py +0 -48
attention.py DELETED
@@ -1,48 +0,0 @@
1
- import torch
2
- import torch.nn as nn
3
- import torch.nn.functional as F
4
-
5
- from .rope import apply_rope
6
-
7
-
8
- class CausalSelfAttention(nn.Module):
9
- """Full multi-head causal self-attention (Section 4.4).
10
-
11
- Deliberately NOT using Grouped Query Attention (GQA) — the doc is explicit
12
- that at this scale, GQA's memory savings are negligible and it can quietly
13
- cost quality. Every head gets its own independent K/V projections.
14
- """
15
-
16
- def __init__(self, hidden_dim: int, n_heads: int, dropout: float = 0.0):
17
- super().__init__()
18
- assert hidden_dim % n_heads == 0
19
- self.n_heads = n_heads
20
- self.head_dim = hidden_dim // n_heads
21
- self.dropout = dropout
22
-
23
- # separate q, k, v projections -- no sharing across heads (full attention)
24
- self.q_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
25
- self.k_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
26
- self.v_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
27
- self.out_proj = nn.Linear(hidden_dim, hidden_dim, bias=False)
28
-
29
- def forward(self, x: torch.Tensor, rope_freqs: torch.Tensor) -> torch.Tensor:
30
- B, T, C = x.shape
31
-
32
- q = self.q_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
33
- k = self.k_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
34
- v = self.v_proj(x).view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
35
-
36
- q = apply_rope(q, rope_freqs[:T])
37
- k = apply_rope(k, rope_freqs[:T])
38
-
39
- # scaled dot-product attention with causal masking (built-in flash-attention
40
- # kernel when running on a CUDA GPU; falls back to a math kernel on CPU)
41
- out = F.scaled_dot_product_attention(
42
- q, k, v,
43
- is_causal=True,
44
- dropout_p=self.dropout if self.training else 0.0,
45
- )
46
-
47
- out = out.transpose(1, 2).contiguous().view(B, T, C)
48
- return self.out_proj(out)