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