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

Method __init__

linear_moe/sequence_modeling/based/based.py:14–45  ·  view source on GitHub ↗
(
        self, 
        config,
        expand_k: float = 1.0,
        expand_v: float = 1.0,
    )

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 1

TaylorFeatureMapClass · 0.90

Tested by

no test coverage detected