MCPcopy Create free account
hub / github.com/OpenSparseLLMs/Linear-MoE / __init__

Method __init__

linear_moe/sequence_modeling/rebased/rebased.py:14–51  ·  view source on GitHub ↗
(
        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,
    )

Source from the content-addressed store, hash-verified

12class 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):

Callers

nothing calls this directly

Calls 1

RebasedFeatureMapClass · 0.90

Tested by

no test coverage detected