Initialize the Attention module. Args: args (ModelArgs): Model configuration parameters. Attributes: n_kv_heads (int): Number of key and value heads. n_local_heads (int): Number of local query heads. n_local_kv_heads (int): N
(self, args: ModelArgs)
| 176 | """Multi-head attention module.""" |
| 177 | |
| 178 | def __init__(self, args: ModelArgs): |
| 179 | """ |
| 180 | Initialize the Attention module. |
| 181 | |
| 182 | Args: |
| 183 | args (ModelArgs): Model configuration parameters. |
| 184 | |
| 185 | Attributes: |
| 186 | n_kv_heads (int): Number of key and value heads. |
| 187 | n_local_heads (int): Number of local query heads. |
| 188 | n_local_kv_heads (int): Number of local key and value heads. |
| 189 | n_rep (int): Number of repetitions for local heads. |
| 190 | head_dim (int): Dimension size of each attention head. |
| 191 | wq (ColumnParallelLinear): Linear transformation for queries. |
| 192 | wk (ColumnParallelLinear): Linear transformation for keys. |
| 193 | wv (ColumnParallelLinear): Linear transformation for values. |
| 194 | wo (RowParallelLinear): Linear transformation for output. |
| 195 | cache_k (torch.Tensor): Cached keys for attention. |
| 196 | cache_v (torch.Tensor): Cached values for attention. |
| 197 | |
| 198 | """ |
| 199 | super().__init__() |
| 200 | self.n_kv_heads = args.n_heads if args.n_kv_heads is None else args.n_kv_heads |
| 201 | model_parallel_size = fs_init.get_model_parallel_world_size() |
| 202 | self.n_local_heads = args.n_heads // model_parallel_size |
| 203 | self.n_local_kv_heads = self.n_kv_heads // model_parallel_size |
| 204 | self.n_rep = self.n_local_heads // self.n_local_kv_heads |
| 205 | self.head_dim = args.dim // args.n_heads |
| 206 | |
| 207 | self.wq = ColumnParallelLinear( |
| 208 | args.dim, |
| 209 | args.n_heads * self.head_dim, |
| 210 | bias=False, |
| 211 | gather_output=False, |
| 212 | init_method=lambda x: x, |
| 213 | ) |
| 214 | self.wk = ColumnParallelLinear( |
| 215 | args.dim, |
| 216 | self.n_kv_heads * self.head_dim, |
| 217 | bias=False, |
| 218 | gather_output=False, |
| 219 | init_method=lambda x: x, |
| 220 | ) |
| 221 | self.wv = ColumnParallelLinear( |
| 222 | args.dim, |
| 223 | self.n_kv_heads * self.head_dim, |
| 224 | bias=False, |
| 225 | gather_output=False, |
| 226 | init_method=lambda x: x, |
| 227 | ) |
| 228 | self.wo = RowParallelLinear( |
| 229 | args.n_heads * self.head_dim, |
| 230 | args.dim, |
| 231 | bias=False, |
| 232 | input_is_parallel=True, |
| 233 | init_method=lambda x: x, |
| 234 | ) |
| 235 |