File size: 6,639 Bytes
3b2d368 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 | import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from torchtune.modules import RotaryPositionalEmbeddings
from torch.nn.attention.flex_attention import flex_attention
from torch.nn.attention import sdpa_kernel, SDPBackend
class ICLAttention(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
# Check if LoRA is enabled for ICL attention
use_lora = getattr(config, 'use_lora_icl_attention', False)
if use_lora:
from .lora import LoRALinear
lora_rank = getattr(config, 'lora_rank', 8)
lora_alpha = getattr(config, 'lora_alpha', 16)
lora_dropout = getattr(config, 'lora_dropout', 0.0)
self.W_q = LoRALinear(
config.embed_dim_phi, config.hidden_dim_f, bias=False,
lora_rank=lora_rank, lora_alpha=lora_alpha, lora_dropout=lora_dropout
)
self.W_k = LoRALinear(
config.embed_dim_phi, config.hidden_dim_f, bias=False,
lora_rank=lora_rank, lora_alpha=lora_alpha, lora_dropout=lora_dropout
)
self.W_v = LoRALinear(
config.embed_dim_f, config.hidden_dim_f, bias=True,
lora_rank=lora_rank, lora_alpha=lora_alpha, lora_dropout=lora_dropout
)
self.W_o = LoRALinear(
config.hidden_dim_f, config.embed_dim_f, bias=True,
lora_rank=lora_rank, lora_alpha=lora_alpha, lora_dropout=lora_dropout
)
else:
self.W_q = nn.Linear(config.embed_dim_phi, config.hidden_dim_f, bias=False)
self.W_k = nn.Linear(config.embed_dim_phi, config.hidden_dim_f, bias=False)
self.W_v = nn.Linear(config.embed_dim_f, config.hidden_dim_f, bias=True)
self.W_o = nn.Linear(config.hidden_dim_f, config.embed_dim_f, bias=True)
self.rotary_embeddings = RotaryPositionalEmbeddings(
config.hidden_dim_f // config.n_heads_f,
max_seq_len=config.max_seq_len + 10
)
self.drop_resid = nn.Dropout(0.1)
def forward(self, q, k, v):
B, S, E = q.shape
q = self.W_q(q).view(B, S, self.config.n_heads_f, self.config.hidden_dim_f // self.config.n_heads_f).transpose(1, 2).contiguous()
k = self.W_k(k).view(B, S, self.config.n_heads_f, self.config.hidden_dim_f // self.config.n_heads_f).transpose(1, 2).contiguous()
v = self.W_v(v).view(B, S, self.config.n_heads_f, self.config.hidden_dim_f // self.config.n_heads_f).transpose(1, 2).contiguous()
q = self.rotary_embeddings(q)
k = self.rotary_embeddings(k)
def _score_mod(scores, b, h, i, j):
keep = (j < i) | ((i == 0) & (j == 0))
return torch.where(keep, scores, torch.full_like(scores, float("-inf")))
with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
attn_output = flex_attention(
q, k, v,
score_mod=_score_mod,
scale=None,
enable_gqa=False
)
attn_output = attn_output.transpose(1, 2).contiguous()
attn_output = attn_output.view(B, S, self.config.hidden_dim_f)
attn_output = self.W_o(attn_output)
attn_output = self.drop_resid(attn_output)
return attn_output
class PhiAttention(nn.Module):
def __init__(self, config):
super().__init__()
self.config = config
# Check if LoRA is enabled for Phi attention
use_lora = getattr(config, 'use_lora_phi_attention', False)
if use_lora:
from .lora import LoRALinear
lora_rank = getattr(config, 'lora_rank', 8)
lora_alpha = getattr(config, 'lora_alpha', 16)
lora_dropout = getattr(config, 'lora_dropout', 0.0)
self.W_q = LoRALinear(
config.embed_dim_phi, config.hidden_dim_phi, bias=False,
lora_rank=lora_rank, lora_alpha=lora_alpha, lora_dropout=lora_dropout
)
self.W_k = LoRALinear(
config.embed_dim_phi, config.hidden_dim_phi, bias=False,
lora_rank=lora_rank, lora_alpha=lora_alpha, lora_dropout=lora_dropout
)
self.W_v = LoRALinear(
config.embed_dim_phi, config.hidden_dim_phi, bias=True,
lora_rank=lora_rank, lora_alpha=lora_alpha, lora_dropout=lora_dropout
)
self.W_o = LoRALinear(
config.hidden_dim_phi, config.embed_dim_phi, bias=True,
lora_rank=lora_rank, lora_alpha=lora_alpha, lora_dropout=lora_dropout
)
else:
self.W_q = nn.Linear(config.embed_dim_phi, config.hidden_dim_phi, bias=False)
self.W_k = nn.Linear(config.embed_dim_phi, config.hidden_dim_phi, bias=False)
self.W_v = nn.Linear(config.embed_dim_phi, config.hidden_dim_phi, bias=True)
self.W_o = nn.Linear(config.hidden_dim_phi, config.embed_dim_phi, bias=True)
self.rotary_embeddings = RotaryPositionalEmbeddings(
config.hidden_dim_phi // config.n_heads_phi,
max_seq_len=config.max_seq_len + 10
)
self.drop_resid = nn.Dropout(0.1)
def forward(self, x):
B, S, E = x.shape
q = self.W_q(x).view(B, S, self.config.n_heads_phi, self.config.hidden_dim_phi // self.config.n_heads_phi).transpose(1, 2).contiguous()
k = self.W_k(x).view(B, S, self.config.n_heads_phi, self.config.hidden_dim_phi // self.config.n_heads_phi).transpose(1, 2).contiguous()
v = self.W_v(x).view(B, S, self.config.n_heads_phi, self.config.hidden_dim_phi // self.config.n_heads_phi).transpose(1, 2).contiguous()
q = self.rotary_embeddings(q)
k = self.rotary_embeddings(k)
def _score_mod(scores, b, h, i, j):
keep = j <= i
return torch.where(keep, scores, torch.full_like(scores, float("-inf")))
with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
attn_output = flex_attention(
q, k, v,
score_mod=_score_mod,
scale=None,
enable_gqa=False
)
attn_output = attn_output.transpose(1, 2).contiguous().view(B, S, self.config.hidden_dim_phi)
attn_output = self.W_o(attn_output)
attn_output = self.drop_resid(attn_output)
return attn_output |