| 164 | |
| 165 | |
| 166 | class QuantLlamaAttentionFused(nn.Module): |
| 167 | def __init__(self, hidden_size, num_heads, qkv_layer, o_proj, dev, args): |
| 168 | super().__init__() |
| 169 | |
| 170 | self.args = args |
| 171 | self.n_local_heads = args.num_attention_heads |
| 172 | self.hidden_size = args.hidden_size |
| 173 | self.num_heads = args.num_attention_heads |
| 174 | self.head_dim = self.hidden_size // self.num_heads |
| 175 | |
| 176 | self.num_key_value_heads = args.num_key_value_heads |
| 177 | self.num_key_value_groups = self.num_heads // self.num_key_value_heads |
| 178 | self.max_position_embeddings = args.max_position_embeddings |
| 179 | # self.rope_theta = args.rope_theta |
| 180 | |
| 181 | self.qkv_proj = qkv_layer |
| 182 | self.o_proj = o_proj |
| 183 | |
| 184 | kv_max_seq_len = min(max_seq_len, args.max_position_embeddings) |
| 185 | |
| 186 | # following fastertransformer definition |
| 187 | self.cache_v = ( |
| 188 | torch.zeros( |
| 189 | ( |
| 190 | max_batch_size, |
| 191 | self.num_key_value_heads, |
| 192 | # args.max_position_embeddings, |
| 193 | kv_max_seq_len, |
| 194 | self.head_dim, |
| 195 | ) |
| 196 | ) |
| 197 | .to(dev) |
| 198 | .half() |
| 199 | ) # added to half |
| 200 | # 8: pack 8 fp16 in FT, if fp32 then use 4 |
| 201 | self.cache_k = ( |
| 202 | torch.zeros( |
| 203 | ( |
| 204 | max_batch_size, |
| 205 | self.num_key_value_heads, |
| 206 | self.head_dim // 8, |
| 207 | # args.max_position_embeddings, |
| 208 | kv_max_seq_len, |
| 209 | 8, |
| 210 | ) |
| 211 | ) |
| 212 | .to(dev) |
| 213 | .half() |
| 214 | ) # added to half |
| 215 | |
| 216 | # dummy |
| 217 | self.rotary_emb = QuantLlamaRotaryEmbedding( |
| 218 | self.head_dim, max_position_embeddings=2048, device="cuda:0" |
| 219 | ) |
| 220 | |
| 221 | def forward( |
| 222 | self, |
| 223 | x: torch.Tensor, |