| 14 | class GLA(MegatronModule): |
| 15 | |
| 16 | def __init__( |
| 17 | self, |
| 18 | config, |
| 19 | expand_k: float = 1.0, |
| 20 | expand_v: float = 1.0, |
| 21 | ): |
| 22 | super().__init__(config) |
| 23 | |
| 24 | self.la_mode = config.la_mode |
| 25 | self.hidden_size = config.hidden_size |
| 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 | |
| 31 | self.la_feature_map = config.la_feature_map |
| 32 | self.la_feature_map_fn = ACT2FN[self.la_feature_map] if self.la_feature_map is not None else None |
| 33 | |
| 34 | self.key_dim = int(config.hidden_size * expand_k) |
| 35 | self.value_dim = int(config.hidden_size * expand_v) |
| 36 | |
| 37 | assert self.la_mode in ['chunk', 'fused_chunk', 'fused_recurrent'], f"Not supported mode `{self.la_mode}`." |
| 38 | assert self.key_dim % self.num_heads == 0, f"key dim must be divisible by num_heads of {self.num_heads}" |
| 39 | assert self.value_dim % self.num_heads == 0, f"value dim must be divisible by num_heads of {self.num_heads}" |
| 40 | |
| 41 | self.head_qk_dim = self.key_dim // self.num_heads |
| 42 | self.head_v_dim = self.value_dim // self.num_heads |
| 43 | |
| 44 | if config.la_output_norm == 'rmsnorm': |
| 45 | self.la_output_norm = RMSNorm(hidden_size=self.head_v_dim, elementwise_affine=config.la_elementwise_affine, eps=config.la_norm_eps) |
| 46 | elif config.la_output_norm == 'identity': |
| 47 | self.la_output_norm = torch.nn.Identity() |
| 48 | else: |
| 49 | raise NotImplementedError(f"Not supported output norm `{self.la_output_norm}`.") |
| 50 | |
| 51 | self.gla_la_gate_logit_normalizer = config.gla_la_gate_logit_normalizer |
| 52 | self.gla_la_clamp_min = config.gla_la_clamp_min |
| 53 | |
| 54 | if self.la_mode == 'chunk': |
| 55 | self._la_impl = chunk_gla |
| 56 | elif self.la_mode == 'fused_chunk': |
| 57 | self._la_impl = fused_chunk_gla |
| 58 | elif self.la_mode == 'fused_recurrent': |
| 59 | self._la_impl = fused_recurrent_gla |
| 60 | |
| 61 | self.apply(self._initialize_weights) |
| 62 | |
| 63 | def _initialize_weights(self, module: torch.nn.Module): |
| 64 | if getattr(module, "_is_hf_initialized", False): |