(
self,
config,
expand_k: float = 1.0,
expand_v: float = 1.0,
)
| 12 | class Based(MegatronModule): |
| 13 | |
| 14 | def __init__( |
| 15 | self, |
| 16 | config, |
| 17 | expand_k: float = 1.0, |
| 18 | expand_v: float = 1.0, |
| 19 | ): |
| 20 | super().__init__(config) |
| 21 | |
| 22 | self.la_mode = config.la_mode |
| 23 | self.hidden_size = config.hidden_size |
| 24 | self.key_dim = int(config.hidden_size * expand_k) |
| 25 | self.value_dim = int(config.hidden_size * expand_v) |
| 26 | self.num_heads = config.num_attention_heads |
| 27 | # num_kv_heads here mains num_query_groups |
| 28 | self.num_kv_heads = config.num_query_groups if config.num_query_groups is not None else config.num_attention_heads |
| 29 | self.num_kv_groups = self.num_heads // self.num_kv_heads |
| 30 | self.head_qk_dim = self.key_dim // self.num_heads |
| 31 | self.head_v_dim = self.value_dim // self.num_heads |
| 32 | self.la_feature_map_fn = TaylorFeatureMap(self.head_qk_dim) |
| 33 | |
| 34 | assert self.la_mode in ['chunk', 'fused_chunk', 'parallel'], f"Not supported mode `{self.la_mode}`." |
| 35 | assert self.key_dim % self.num_heads == 0, f"key dim must be divisible by num_heads of {self.num_heads}" |
| 36 | assert self.value_dim % self.num_heads == 0, f"value dim must be divisible by num_heads of {self.num_heads}" |
| 37 | |
| 38 | if self.la_mode == 'chunk': |
| 39 | self._la_impl = chunk_linear_attn |
| 40 | elif self.la_mode == 'fused_chunk': |
| 41 | self._la_impl = fused_chunk_linear_attn |
| 42 | elif self.la_mode == 'parallel': |
| 43 | self._la_impl = parallel_based |
| 44 | |
| 45 | self.apply(self._initialize_weights) |
| 46 | |
| 47 | def _initialize_weights(self, module: torch.nn.Module): |
| 48 | if getattr(module, "_is_hf_initialized", False): |
nothing calls this directly
no test coverage detected