Multi-headed attention from 'Attention Is All You Need' paper
| 77 | |
| 78 | |
| 79 | class QuantLlamaAttention(nn.Module): |
| 80 | """Multi-headed attention from 'Attention Is All You Need' paper""" |
| 81 | |
| 82 | def __init__(self, hidden_size, num_heads, qkv_proj, o_proj, dev): |
| 83 | super().__init__() |
| 84 | self.hidden_size = hidden_size |
| 85 | self.num_heads = num_heads |
| 86 | self.head_dim = hidden_size // num_heads |
| 87 | |
| 88 | if (self.head_dim * num_heads) != self.hidden_size: |
| 89 | raise ValueError( |
| 90 | f"hidden_size must be divisible by num_heads (got `hidden_size`: {self.hidden_size}" |
| 91 | f" and `num_heads`: {num_heads})." |
| 92 | ) |
| 93 | self.qkv_proj = qkv_proj |
| 94 | self.o_proj = o_proj |
| 95 | self.rotary_emb = QuantLlamaRotaryEmbedding( |
| 96 | self.head_dim, max_position_embeddings=2048, device=dev |
| 97 | ) |
| 98 | |
| 99 | def forward( |
| 100 | self, |
| 101 | hidden_states, |
| 102 | past_key_value=None, |
| 103 | attention_mask=None, |
| 104 | position_ids=None, |
| 105 | output_attentions=False, |
| 106 | use_cache=False, |
| 107 | ): |
| 108 | """Input shape: Batch x Time x Channel""" |
| 109 | |
| 110 | bsz, q_len, _ = hidden_states.size() |
| 111 | |
| 112 | qkv_states = self.qkv_proj(hidden_states) |
| 113 | qkv_states = qkv_states.view(bsz, q_len, 3, self.num_heads, self.head_dim) |
| 114 | |
| 115 | # This updates the query and key states in-place, saving VRAM. |
| 116 | query_states, key_states, value_states = torch.split(qkv_states, 1, dim=2) |
| 117 | query_states, key_states = self.rotary_emb( |
| 118 | query_states, key_states, position_ids |
| 119 | ) |
| 120 | |
| 121 | del qkv_states |
| 122 | query_states = query_states.view( |
| 123 | bsz, q_len, self.num_heads, self.head_dim |
| 124 | ).transpose(1, 2) |
| 125 | key_states = key_states.view( |
| 126 | bsz, q_len, self.num_heads, self.head_dim |
| 127 | ).transpose(1, 2) |
| 128 | value_states = value_states.view( |
| 129 | bsz, q_len, self.num_heads, self.head_dim |
| 130 | ).transpose(1, 2) |
| 131 | |
| 132 | is_causal = past_key_value is None |
| 133 | |
| 134 | kv_seq_len = q_len |
| 135 | if past_key_value is not None: |
| 136 | kv_seq_len += past_key_value[0].shape[-2] |