devoppro commited on
Commit
af0d032
·
verified ·
1 Parent(s): 1114313

Upload 3 files

Browse files
config.json CHANGED
@@ -2,9 +2,13 @@
2
  "architectures": [
3
  "ModernLLMForCausalLM"
4
  ],
5
- "bos_token_id": 1,
 
 
 
 
6
  "dtype": "float32",
7
- "eos_token_id": 2,
8
  "hidden_size": 768,
9
  "intermediate_size": 2048,
10
  "max_position_embeddings": 2048,
@@ -12,7 +16,7 @@
12
  "num_attention_heads": 12,
13
  "num_hidden_layers": 12,
14
  "num_key_value_heads": 4,
15
- "pad_token_id": 0,
16
  "rms_norm_eps": 1e-06,
17
  "rope_theta": 1000000.0,
18
  "transformers_version": "5.15.1",
 
2
  "architectures": [
3
  "ModernLLMForCausalLM"
4
  ],
5
+ "auto_map": {
6
+ "AutoConfig": "configuration_modern_llm.ModernLLMConfig",
7
+ "AutoModelForCausalLM": "modeling_modern_llm.ModernLLMForCausalLM"
8
+ },
9
+ "bos_token_id": 151643,
10
  "dtype": "float32",
11
+ "eos_token_id": 151643,
12
  "hidden_size": 768,
13
  "intermediate_size": 2048,
14
  "max_position_embeddings": 2048,
 
16
  "num_attention_heads": 12,
17
  "num_hidden_layers": 12,
18
  "num_key_value_heads": 4,
19
+ "pad_token_id": 151643,
20
  "rms_norm_eps": 1e-06,
21
  "rope_theta": 1000000.0,
22
  "transformers_version": "5.15.1",
configuration_modern_llm.py ADDED
@@ -0,0 +1,37 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from transformers import PretrainedConfig
2
+
3
+
4
+ class ModernLLMConfig(PretrainedConfig):
5
+ model_type = "modern_llm"
6
+
7
+ def __init__(
8
+ self,
9
+ vocab_size: int = 151936,
10
+ hidden_size: int = 768,
11
+ intermediate_size: int = 2048,
12
+ num_hidden_layers: int = 12,
13
+ num_attention_heads: int = 12,
14
+ num_key_value_heads: int = 4, # Grouped-Query Attention (GQA)
15
+ max_position_embeddings: int = 2048,
16
+ rms_norm_eps: float = 1e-6,
17
+ rope_theta: float = 1000000.0,
18
+ pad_token_id: int = 0,
19
+ bos_token_id: int = 1,
20
+ eos_token_id: int = 2,
21
+ **kwargs,
22
+ ):
23
+ self.vocab_size = vocab_size
24
+ self.hidden_size = hidden_size
25
+ self.intermediate_size = intermediate_size
26
+ self.num_hidden_layers = num_hidden_layers
27
+ self.num_attention_heads = num_attention_heads
28
+ self.num_key_value_heads = num_key_value_heads
29
+ self.max_position_embeddings = max_position_embeddings
30
+ self.rms_norm_eps = rms_norm_eps
31
+ self.rope_theta = rope_theta
32
+ super().__init__(
33
+ pad_token_id=pad_token_id,
34
+ bos_token_id=bos_token_id,
35
+ eos_token_id=eos_token_id,
36
+ **kwargs,
37
+ )
modeling_modern_llm.py ADDED
@@ -0,0 +1,137 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Optional
2
+
3
+ import torch
4
+ import torch.nn as nn
5
+ import torch.nn.functional as F
6
+ from transformers import PreTrainedModel
7
+
8
+ from .configuration_modern_llm import ModernLLMConfig
9
+
10
+
11
+ class RMSNorm(nn.Module):
12
+ def __init__(self, dim: int, eps: float = 1e-6):
13
+ super().__init__()
14
+ self.eps = eps
15
+ self.weight = nn.Parameter(torch.ones(dim))
16
+
17
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
18
+ variance = x.pow(2).mean(-1, keepdim=True)
19
+ return x * torch.rsqrt(variance + self.eps) * self.weight
20
+
21
+
22
+ class RotaryEmbedding(nn.Module):
23
+ def __init__(self, dim: int, max_position_embeddings: int = 2048, base: float = 1000000.0):
24
+ super().__init__()
25
+ inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))
26
+ self.register_buffer("inv_freq", inv_freq, persistent=False)
27
+
28
+ def forward(self, x: torch.Tensor, seq_len: int):
29
+ t = torch.arange(seq_len, device=x.device, dtype=self.inv_freq.dtype)
30
+ freqs = torch.outer(t, self.inv_freq)
31
+ emb = torch.cat((freqs, freqs), dim=-1)
32
+ return emb.cos(), emb.sin()
33
+
34
+
35
+ def rotate_half(x: torch.Tensor) -> torch.Tensor:
36
+ x1 = x[..., : x.shape[-1] // 2]
37
+ x2 = x[..., x.shape[-1] // 2 :]
38
+ return torch.cat((-x2, x1), dim=-1)
39
+
40
+
41
+ def apply_rotary_pos_emb(q, k, cos, sin):
42
+ cos = cos.unsqueeze(0).unsqueeze(2)
43
+ sin = sin.unsqueeze(0).unsqueeze(2)
44
+ q_embed = (q * cos) + (rotate_half(q) * sin)
45
+ k_embed = (k * cos) + (rotate_half(k) * sin)
46
+ return q_embed, k_embed
47
+
48
+
49
+ class SwiGLU(nn.Module):
50
+ def __init__(self, config: ModernLLMConfig):
51
+ super().__init__()
52
+ self.gate_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)
53
+ self.up_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)
54
+ self.down_proj = nn.Linear(config.intermediate_size, config.hidden_size, bias=False)
55
+
56
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
57
+ return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
58
+
59
+
60
+ class GroupedQueryAttention(nn.Module):
61
+ def __init__(self, config: ModernLLMConfig):
62
+ super().__init__()
63
+ self.num_heads = config.num_attention_heads
64
+ self.head_dim = config.hidden_size // config.num_attention_heads
65
+ self.num_kv_heads = config.num_key_value_heads
66
+ self.num_kv_groups = self.num_heads // self.num_kv_heads
67
+
68
+ self.q_proj = nn.Linear(config.hidden_size, self.num_heads * self.head_dim, bias=False)
69
+ self.k_proj = nn.Linear(config.hidden_size, self.num_kv_heads * self.head_dim, bias=False)
70
+ self.v_proj = nn.Linear(config.hidden_size, self.num_kv_heads * self.head_dim, bias=False)
71
+ self.o_proj = nn.Linear(self.num_heads * self.head_dim, config.hidden_size, bias=False)
72
+
73
+ def forward(self, x: torch.Tensor, rot_cos: torch.Tensor, rot_sin: torch.Tensor) -> torch.Tensor:
74
+ batch_size, seq_len, _ = x.shape
75
+ q = self.q_proj(x).view(batch_size, seq_len, self.num_heads, self.head_dim)
76
+ k = self.k_proj(x).view(batch_size, seq_len, self.num_kv_heads, self.head_dim)
77
+ v = self.v_proj(x).view(batch_size, seq_len, self.num_kv_heads, self.head_dim)
78
+
79
+ q, k = apply_rotary_pos_emb(q, k, rot_cos, rot_sin)
80
+
81
+ k = k.repeat_interleave(self.num_kv_groups, dim=2)
82
+ v = v.repeat_interleave(self.num_kv_groups, dim=2)
83
+
84
+ q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)
85
+ out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
86
+ out = out.transpose(1, 2).contiguous().view(batch_size, seq_len, -1)
87
+ return self.o_proj(out)
88
+
89
+
90
+ class TransformerBlock(nn.Module):
91
+ def __init__(self, config: ModernLLMConfig):
92
+ super().__init__()
93
+ self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
94
+ self.self_attn = GroupedQueryAttention(config)
95
+ self.post_attention_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
96
+ self.mlp = SwiGLU(config)
97
+
98
+ def forward(self, x: torch.Tensor, rot_cos: torch.Tensor, rot_sin: torch.Tensor) -> torch.Tensor:
99
+ x = x + self.self_attn(self.input_layernorm(x), rot_cos, rot_sin)
100
+ x = x + self.mlp(self.post_attention_layernorm(x))
101
+ return x
102
+
103
+
104
+ class ModernLLMForCausalLM(PreTrainedModel):
105
+ config_class = ModernLLMConfig
106
+
107
+ def __init__(self, config: ModernLLMConfig):
108
+ super().__init__(config)
109
+ self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size)
110
+ self.layers = nn.ModuleList([TransformerBlock(config) for _ in range(config.num_hidden_layers)])
111
+ self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
112
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
113
+ self.rotary_emb = RotaryEmbedding(
114
+ config.hidden_size // config.num_attention_heads,
115
+ config.max_position_embeddings,
116
+ config.rope_theta,
117
+ )
118
+ self.post_init()
119
+
120
+ def forward(self, input_ids: torch.LongTensor, labels: Optional[torch.LongTensor] = None, **kwargs):
121
+ _, seq_len = input_ids.shape
122
+ x = self.embed_tokens(input_ids)
123
+ cos, sin = self.rotary_emb(x, seq_len)
124
+
125
+ for layer in self.layers:
126
+ x = layer(x, cos, sin)
127
+
128
+ x = self.norm(x)
129
+ logits = self.lm_head(x)
130
+
131
+ loss = None
132
+ if labels is not None:
133
+ shift_logits = logits[..., :-1, :].contiguous()
134
+ shift_labels = labels[..., 1:].contiguous()
135
+ loss = F.cross_entropy(shift_logits.view(-1, self.config.vocab_size), shift_labels.view(-1))
136
+
137
+ return {"loss": loss, "logits": logits}