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

Method __init__

linear_moe/sequence_modeling/gla/gla.py:16–61  ·  view source on GitHub ↗
(
        self, 
        config,
        expand_k: float = 1.0,
        expand_v: float = 1.0,
    )

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 1

RMSNormClass · 0.90

Tested by

no test coverage detected