Banaxi-Tech commited on
Commit
68d4e14
·
verified ·
1 Parent(s): 09c7d0f

Upload 11 files

Browse files
.gitattributes CHANGED
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ banner.png filter=lfs diff=lfs merge=lfs -text
README.md CHANGED
@@ -1,3 +1,175 @@
1
  ---
2
  license: apache-2.0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3
  ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  ---
2
  license: apache-2.0
3
+ language:
4
+ - en
5
+ library_name: transformers
6
+ pipeline_tag: text-generation
7
+ datasets:
8
+ - HuggingFaceFW/fineweb-edu
9
+ - mlfoundations/dclm-baseline-1.0
10
+ - HuggingFaceTB/smollm-corpus
11
+ - HuggingFaceTB/finemath
12
+ tags:
13
+ - causal-lm
14
+ - language-model
15
+ - base-model
16
+ - small-language-model
17
+ - bananamind
18
+ - bananamind2
19
+ - bananamind2-nano
20
+ - fineweb-edu
21
+ - dclm
22
+ - cosmopedia-v2
23
+ - finemath
24
+ - digit-tokenizer
25
+ - pytorch
26
+ - safetensors
27
+ - custom-code
28
+ - trust-remote-code
29
+ - custom-architecture
30
  ---
31
+
32
+ # BananaMind-2-Nano
33
+ ![Banner](banner.png)
34
+
35
+ BananaMind-2-Nano is a compact decoder-only causal language model trained from scratch by BananaMind on a 30B-token curriculum.
36
+
37
+ The model has **9,968,128 parameters**, a **4,096-token context window**, and a custom **8k-token digit-aware byte-level BPE tokenizer**.
38
+
39
+ ## Model Details
40
+
41
+ | Field | Value |
42
+ |---|---:|
43
+ | Parameters | 9,968,128 |
44
+ | Architecture | BananaMind2Nano decoder-only Transformer |
45
+ | Layers | 10 |
46
+ | Hidden size | 256 |
47
+ | Intermediate size | 768 |
48
+ | Attention heads | 4 |
49
+ | KV heads | 2 |
50
+ | Head dim | 64 |
51
+ | Attention style | Grouped-query attention with QK norm |
52
+ | MLP | SwiGLU |
53
+ | Position embeddings | RoPE |
54
+ | RoPE theta | 100,000 |
55
+ | Normalization | RMSNorm |
56
+ | RMSNorm epsilon | 1e-06 |
57
+ | Vocabulary size | 8,192 |
58
+ | Context length | 4,096 |
59
+ | Embeddings | Tied input/output embeddings |
60
+ | Weight format | safetensors |
61
+ | HF architecture | `BananaMind2NanoForCausalLM` |
62
+ | HF model type | `bananamind2_nano` |
63
+ | Final checkpoint | `runs/bananamind2-nano/final.pt` |
64
+ | Final training step | 55,485 |
65
+ | Tokens seen | 29,999,726,592 |
66
+
67
+ ## Tokenizer
68
+
69
+ BananaMind-2-Nano uses the same custom 8k byte-level BPE tokenizer as BananaMind-2-Mini. Digits are kept as separate tokens so numbers do not collapse into large number tokens.
70
+
71
+ | Special token | ID |
72
+ |---|---:|
73
+ | `<|pad|>` | 0 |
74
+ | `<|bos|>` | 1 |
75
+ | `<|eos|>` | 2 |
76
+ | `<|unk|>` | 3 |
77
+
78
+ ## Training Data
79
+
80
+ | Dataset | Target Tokens | Share |
81
+ |---|---:|---:|
82
+ | FineWeb-Edu | 16.5B | 55% |
83
+ | DCLM | 9.0B | 30% |
84
+ | Cosmopedia-v2 | 3.0B | 10% |
85
+ | FineMath-4+ | 1.5B | 5% |
86
+ | Total | 30.0B | 100% |
87
+
88
+ The run used a progressive curriculum, beginning web-heavy and gradually increasing synthetic textbook and mathematics data.
89
+
90
+ ## Training Setup
91
+
92
+ | Field | Value |
93
+ |---|---:|
94
+ | Sequence length | 4,096 |
95
+ | Micro batch | 12 |
96
+ | Gradient accumulation | 11 |
97
+ | Effective batch | 132 sequences |
98
+ | Tokens per optimizer step | 540,672 |
99
+ | Final optimizer step | 55,485 |
100
+ | Optimizer | AdamW |
101
+ | Betas | 0.9, 0.95 |
102
+ | Peak learning rate | 0.003 |
103
+ | Warmup steps | 1,750 |
104
+ | LR schedule | Warmup-stable-decay with cosine decay |
105
+ | Weight decay | 0.1, then 0.01 after 12B tokens |
106
+ | Gradient clipping | 1 |
107
+ | Z-loss coefficient | 1e-4 until 12B tokens, then off |
108
+ | Compile | PyTorch compile enabled |
109
+ | Seed | 1337 |
110
+
111
+ ## Evaluation
112
+
113
+ These are self-reported scores produced with `lm_eval`. Scores may vary slightly depending on the evaluation harness version, runtime settings, dtype, and environment.
114
+
115
+ All task scores use `acc_norm,none`. The average is the mean of ARC Easy, PIQA, ARC Challenge, and HellaSwag.
116
+
117
+ | Benchmark | Score | Metric |
118
+ |---|---:|---|
119
+ | **Average** | **35.77** | `mean` |
120
+ | ARC Easy | 36.20 | `acc_norm,none` |
121
+ | PIQA | 55.98 | `acc_norm,none` |
122
+ | ARC Challenge | 23.38 | `acc_norm,none` |
123
+ | HellaSwag | 27.50 | `acc_norm,none` |
124
+
125
+ The unrounded average is `0.357659`. Available unrounded task results are ARC Easy `0.361953`, PIQA `0.559848`, and ARC Challenge `0.233788`.
126
+
127
+ ## Usage
128
+
129
+ This model uses custom architecture code, so load it with `trust_remote_code=True`.
130
+
131
+ ```bash
132
+ pip install -U transformers safetensors torch
133
+ ```
134
+
135
+ ```python
136
+ import torch
137
+ from transformers import AutoModelForCausalLM, AutoTokenizer
138
+
139
+ model_id = "BananaMind/BananaMind-2-Nano"
140
+ tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
141
+
142
+ device = "cuda" if torch.cuda.is_available() else "cpu"
143
+ dtype = (
144
+ torch.bfloat16
145
+ if torch.cuda.is_available() and torch.cuda.is_bf16_supported()
146
+ else torch.float32
147
+ )
148
+
149
+ model = AutoModelForCausalLM.from_pretrained(
150
+ model_id,
151
+ trust_remote_code=True,
152
+ torch_dtype=dtype,
153
+ ).to(device).eval()
154
+
155
+ prompt = "The color of the sky is"
156
+ inputs = tokenizer(prompt, return_tensors="pt").to(device)
157
+
158
+ with torch.no_grad():
159
+ output = model.generate(
160
+ **inputs,
161
+ max_new_tokens=96,
162
+ do_sample=True,
163
+ temperature=0.7,
164
+ top_p=0.9,
165
+ repetition_penalty=1.1,
166
+ pad_token_id=tokenizer.eos_token_id,
167
+ eos_token_id=tokenizer.eos_token_id,
168
+ )
169
+
170
+ print(tokenizer.decode(output[0], skip_special_tokens=True))
171
+ ```
172
+
173
+ ## License
174
+
175
+ Apache 2.0
banner.png ADDED

Git LFS Details

  • SHA256: 57e0da52e4912a1315b395ce3be4cc7617bde09fe008b3e4550ee71e494a5da6
  • Pointer size: 132 Bytes
  • Size of remote file: 1.38 MB
checkpoint_metadata.json ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "source_checkpoint": "runs/bananamind2-nano/final.pt",
3
+ "step": 55485,
4
+ "tokens_seen": 29999726592,
5
+ "parameters": 9968128,
6
+ "repo_id": "BananaMind/BananaMind-2-Nano",
7
+ "format": "huggingface_transformers_remote_code"
8
+ }
config.json ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "vocab_size": 8192,
3
+ "hidden_size": 256,
4
+ "num_hidden_layers": 10,
5
+ "num_attention_heads": 4,
6
+ "num_key_value_heads": 2,
7
+ "head_dim": 64,
8
+ "intermediate_size": 768,
9
+ "max_position_embeddings": 4096,
10
+ "rope_theta": 100000.0,
11
+ "rms_norm_eps": 1e-06,
12
+ "tie_word_embeddings": true,
13
+ "model_type": "bananamind2_nano",
14
+ "architectures": [
15
+ "BananaMind2NanoForCausalLM"
16
+ ],
17
+ "auto_map": {
18
+ "AutoConfig": "configuration_bananamind2nano.BananaMind2NanoConfig",
19
+ "AutoModelForCausalLM": "modeling_bananamind2nano.BananaMind2NanoForCausalLM"
20
+ },
21
+ "torch_dtype": "float32",
22
+ "transformers_version": "5.7.0",
23
+ "bos_token_id": 1,
24
+ "eos_token_id": 2,
25
+ "pad_token_id": 0,
26
+ "unk_token_id": 3,
27
+ "use_cache": true,
28
+ "z_loss_coeff": 0.0,
29
+ "_name_or_path": "BananaMind/BananaMind-2-Nano"
30
+ }
configuration_bananamind2nano.py ADDED
@@ -0,0 +1,48 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """BananaMind 2 Nano configuration."""
2
+ from transformers import PretrainedConfig
3
+
4
+
5
+ class BananaMind2NanoConfig(PretrainedConfig):
6
+ model_type = "bananamind2_nano"
7
+
8
+ def __init__(
9
+ self,
10
+ vocab_size=8192,
11
+ hidden_size=256,
12
+ num_hidden_layers=10,
13
+ num_attention_heads=4,
14
+ num_key_value_heads=2,
15
+ head_dim=64,
16
+ intermediate_size=768,
17
+ max_position_embeddings=4096,
18
+ rope_theta=100000.0,
19
+ rms_norm_eps=1e-6,
20
+ tie_word_embeddings=True,
21
+ use_cache=True,
22
+ z_loss_coeff=0.0,
23
+ bos_token_id=1,
24
+ eos_token_id=2,
25
+ pad_token_id=0,
26
+ unk_token_id=3,
27
+ **kwargs,
28
+ ):
29
+ self.vocab_size = vocab_size
30
+ self.hidden_size = hidden_size
31
+ self.num_hidden_layers = num_hidden_layers
32
+ self.num_attention_heads = num_attention_heads
33
+ self.num_key_value_heads = num_key_value_heads
34
+ self.head_dim = head_dim
35
+ self.intermediate_size = intermediate_size
36
+ self.max_position_embeddings = max_position_embeddings
37
+ self.rope_theta = rope_theta
38
+ self.rms_norm_eps = rms_norm_eps
39
+ self.use_cache = use_cache
40
+ self.z_loss_coeff = z_loss_coeff
41
+ super().__init__(
42
+ tie_word_embeddings=tie_word_embeddings,
43
+ bos_token_id=bos_token_id,
44
+ eos_token_id=eos_token_id,
45
+ pad_token_id=pad_token_id,
46
+ unk_token_id=unk_token_id,
47
+ **kwargs,
48
+ )
generation_config.json ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ {
2
+ "bos_token_id": 1,
3
+ "eos_token_id": 2,
4
+ "pad_token_id": 0,
5
+ "use_cache": true,
6
+ "transformers_version": "5.7.0"
7
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:80ebf54a809db1cdf7689374f454ba80c329366a99be7f662943d3249fe37f92
3
+ size 48272696
modeling_bananamind2nano.py ADDED
@@ -0,0 +1,307 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """BananaMind 2 Nano implementation for Hugging Face Transformers."""
2
+ import math
3
+ from typing import Optional
4
+
5
+ import torch
6
+ import torch.nn as nn
7
+ import torch.nn.functional as F
8
+ from transformers import PreTrainedModel
9
+ from transformers.cache_utils import Cache, DynamicCache
10
+ from transformers.generation.utils import GenerationMixin
11
+ from transformers.modeling_outputs import CausalLMOutputWithPast
12
+
13
+ from .configuration_bananamind2nano import BananaMind2NanoConfig
14
+
15
+
16
+ class RMSNorm(nn.Module):
17
+ def __init__(self, dim, eps=1e-6):
18
+ super().__init__()
19
+ self.eps = eps
20
+ self.weight = nn.Parameter(torch.ones(dim))
21
+
22
+ def forward(self, x):
23
+ x_float = x.float()
24
+ rms = torch.rsqrt(x_float.pow(2).mean(-1, keepdim=True) + self.eps)
25
+ return (x_float * rms * self.weight.float()).type_as(x)
26
+
27
+
28
+ def build_rope_inv_freq(head_dim, theta=100000.0):
29
+ return 1.0 / (theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim))
30
+
31
+
32
+ def precompute_freqs_cis(head_dim, seq_len, theta=100000.0):
33
+ freqs = build_rope_inv_freq(head_dim, theta)
34
+ positions = torch.arange(seq_len, dtype=torch.float32)
35
+ freqs = torch.outer(positions, freqs)
36
+ return torch.polar(torch.ones_like(freqs), freqs)
37
+
38
+
39
+ def apply_rotary_emb(q, k, freqs_cis):
40
+ q_complex = torch.view_as_complex(q.float().reshape(*q.shape[:-1], -1, 2))
41
+ k_complex = torch.view_as_complex(k.float().reshape(*k.shape[:-1], -1, 2))
42
+ freqs_cis = freqs_cis.unsqueeze(0).unsqueeze(0)
43
+ q_out = torch.view_as_real(q_complex * freqs_cis).flatten(-2)
44
+ k_out = torch.view_as_real(k_complex * freqs_cis).flatten(-2)
45
+ return q_out.type_as(q), k_out.type_as(k)
46
+
47
+
48
+ class BananaMind2NanoAttention(nn.Module):
49
+ def __init__(self, config, layer_idx):
50
+ super().__init__()
51
+ self.layer_idx = layer_idx
52
+ self.n_head = config.num_attention_heads
53
+ self.n_kv_heads = config.num_key_value_heads
54
+ self.head_dim = config.head_dim
55
+ self.n_rep = self.n_head // self.n_kv_heads
56
+
57
+ self.q_proj = nn.Linear(config.hidden_size, self.n_head * self.head_dim, bias=False)
58
+ self.k_proj = nn.Linear(config.hidden_size, self.n_kv_heads * self.head_dim, bias=False)
59
+ self.v_proj = nn.Linear(config.hidden_size, self.n_kv_heads * self.head_dim, bias=False)
60
+ self.o_proj = nn.Linear(self.n_head * self.head_dim, config.hidden_size, bias=False)
61
+ self.q_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps)
62
+ self.k_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps)
63
+
64
+ def forward(
65
+ self,
66
+ x,
67
+ freqs_cis,
68
+ attention_mask=None,
69
+ past_key_values=None,
70
+ use_cache=False,
71
+ ):
72
+ batch_size, seq_len, _ = x.shape
73
+ q = self.q_proj(x).view(
74
+ batch_size,
75
+ seq_len,
76
+ self.n_head,
77
+ self.head_dim,
78
+ ).transpose(1, 2)
79
+ k = self.k_proj(x).view(
80
+ batch_size,
81
+ seq_len,
82
+ self.n_kv_heads,
83
+ self.head_dim,
84
+ ).transpose(1, 2)
85
+ v = self.v_proj(x).view(
86
+ batch_size,
87
+ seq_len,
88
+ self.n_kv_heads,
89
+ self.head_dim,
90
+ ).transpose(1, 2)
91
+
92
+ q = self.q_norm(q)
93
+ k = self.k_norm(k)
94
+ q, k = apply_rotary_emb(q, k, freqs_cis)
95
+
96
+ past_length = 0
97
+ if use_cache and past_key_values is not None:
98
+ past_length = past_key_values.get_seq_length(self.layer_idx)
99
+ k, v = past_key_values.update(k, v, self.layer_idx)
100
+
101
+ kv_len = k.size(-2)
102
+ k = k.unsqueeze(2).expand(
103
+ batch_size,
104
+ self.n_kv_heads,
105
+ self.n_rep,
106
+ kv_len,
107
+ self.head_dim,
108
+ ).reshape(batch_size, self.n_head, kv_len, self.head_dim)
109
+ v = v.unsqueeze(2).expand(
110
+ batch_size,
111
+ self.n_kv_heads,
112
+ self.n_rep,
113
+ kv_len,
114
+ self.head_dim,
115
+ ).reshape(batch_size, self.n_head, kv_len, self.head_dim)
116
+
117
+ attn_mask = None
118
+ is_causal = past_length == 0 and attention_mask is None
119
+ if not is_causal:
120
+ query_positions = past_length + torch.arange(seq_len, device=x.device)
121
+ key_positions = torch.arange(kv_len, device=x.device)
122
+ causal = key_positions.unsqueeze(0) <= query_positions.unsqueeze(1)
123
+ attn_mask = causal[None, None, :, :]
124
+ if attention_mask is not None:
125
+ key_padding = attention_mask.to(torch.bool)
126
+ if key_padding.size(-1) < kv_len:
127
+ cached_padding = torch.ones(
128
+ key_padding.size(0),
129
+ kv_len - key_padding.size(-1),
130
+ dtype=torch.bool,
131
+ device=key_padding.device,
132
+ )
133
+ key_padding = torch.cat((cached_padding, key_padding), dim=-1)
134
+ else:
135
+ key_padding = key_padding[:, -kv_len:]
136
+ attn_mask = attn_mask & key_padding[:, None, None, :]
137
+ is_causal = False
138
+
139
+ y = F.scaled_dot_product_attention(
140
+ q,
141
+ k,
142
+ v,
143
+ attn_mask=attn_mask,
144
+ is_causal=is_causal,
145
+ )
146
+ y = y.transpose(1, 2).contiguous().view(
147
+ batch_size,
148
+ seq_len,
149
+ self.n_head * self.head_dim,
150
+ )
151
+ return self.o_proj(y)
152
+
153
+
154
+ class BananaMind2NanoSwiGLUMLP(nn.Module):
155
+ def __init__(self, config):
156
+ super().__init__()
157
+ self.w_gate = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)
158
+ self.w_up = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)
159
+ self.w_down = nn.Linear(config.intermediate_size, config.hidden_size, bias=False)
160
+
161
+ def forward(self, x):
162
+ return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x))
163
+
164
+
165
+ class BananaMind2NanoBlock(nn.Module):
166
+ def __init__(self, config, layer_idx):
167
+ super().__init__()
168
+ self.ln_1 = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
169
+ self.attn = BananaMind2NanoAttention(config, layer_idx)
170
+ self.ln_2 = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
171
+ self.mlp = BananaMind2NanoSwiGLUMLP(config)
172
+
173
+ def forward(
174
+ self,
175
+ x,
176
+ freqs_cis,
177
+ attention_mask=None,
178
+ past_key_values=None,
179
+ use_cache=False,
180
+ ):
181
+ x = x + self.attn(
182
+ self.ln_1(x),
183
+ freqs_cis,
184
+ attention_mask=attention_mask,
185
+ past_key_values=past_key_values,
186
+ use_cache=use_cache,
187
+ )
188
+ return x + self.mlp(self.ln_2(x))
189
+
190
+
191
+ class BananaMind2NanoPreTrainedModel(PreTrainedModel):
192
+ config_class = BananaMind2NanoConfig
193
+ base_model_prefix = "transformer"
194
+ supports_gradient_checkpointing = False
195
+
196
+ def _init_weights(self, module):
197
+ std = 0.02
198
+ if hasattr(module, "NANOGPT_SCALE_INIT"):
199
+ std *= 2 * self.config.num_hidden_layers ** -0.5
200
+ if isinstance(module, nn.Linear):
201
+ nn.init.normal_(module.weight, mean=0.0, std=std)
202
+ elif isinstance(module, nn.Embedding):
203
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
204
+
205
+
206
+ class BananaMind2NanoForCausalLM(BananaMind2NanoPreTrainedModel, GenerationMixin):
207
+ _tied_weights_keys = {"lm_head.weight": "transformer.wte.weight"}
208
+
209
+ def __init__(self, config):
210
+ super().__init__(config)
211
+ self.config = config
212
+ self.transformer = nn.ModuleDict(
213
+ {
214
+ "wte": nn.Embedding(config.vocab_size, config.hidden_size),
215
+ "h": nn.ModuleList(
216
+ [
217
+ BananaMind2NanoBlock(config, layer_idx)
218
+ for layer_idx in range(config.num_hidden_layers)
219
+ ]
220
+ ),
221
+ "ln_f": RMSNorm(config.hidden_size, eps=config.rms_norm_eps),
222
+ }
223
+ )
224
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
225
+ if config.tie_word_embeddings:
226
+ self.lm_head.weight = self.transformer["wte"].weight
227
+ self._embd_scale = math.sqrt(config.hidden_size)
228
+ self._freqs_cis_cache = None
229
+ self.post_init()
230
+
231
+ def get_input_embeddings(self):
232
+ return self.transformer["wte"]
233
+
234
+ def set_input_embeddings(self, value):
235
+ self.transformer["wte"] = value
236
+
237
+ def get_output_embeddings(self):
238
+ return self.lm_head
239
+
240
+ def set_output_embeddings(self, new_embeddings):
241
+ self.lm_head = new_embeddings
242
+
243
+ def _get_freqs_cis(self, seq_len, device):
244
+ cache = self._freqs_cis_cache
245
+ if cache is None or cache.device != device or cache.size(0) < seq_len:
246
+ cache = precompute_freqs_cis(
247
+ self.config.head_dim,
248
+ seq_len,
249
+ self.config.rope_theta,
250
+ ).to(device)
251
+ self._freqs_cis_cache = cache
252
+ return cache[:seq_len]
253
+
254
+ def forward(
255
+ self,
256
+ input_ids,
257
+ attention_mask=None,
258
+ labels=None,
259
+ past_key_values: Optional[Cache] = None,
260
+ use_cache=None,
261
+ **kwargs,
262
+ ):
263
+ _, seq_len = input_ids.shape
264
+ if use_cache is None:
265
+ use_cache = self.config.use_cache and labels is None
266
+ if use_cache and past_key_values is None:
267
+ past_key_values = DynamicCache(config=self.config)
268
+
269
+ past_length = past_key_values.get_seq_length() if use_cache else 0
270
+ total_length = past_length + seq_len
271
+ if total_length > self.config.max_position_embeddings:
272
+ raise ValueError(
273
+ f"Sequence length {total_length} exceeds the configured maximum "
274
+ f"of {self.config.max_position_embeddings}"
275
+ )
276
+
277
+ x = self.transformer["wte"](input_ids) * self._embd_scale
278
+ freqs_cis = self._get_freqs_cis(total_length, input_ids.device)[past_length:]
279
+
280
+ for block in self.transformer["h"]:
281
+ x = block(
282
+ x,
283
+ freqs_cis,
284
+ attention_mask=attention_mask,
285
+ past_key_values=past_key_values,
286
+ use_cache=use_cache,
287
+ )
288
+
289
+ x = self.transformer["ln_f"](x)
290
+ logits = self.lm_head(x)
291
+
292
+ loss = None
293
+ if labels is not None:
294
+ shift_logits = logits[..., :-1, :].contiguous()
295
+ shift_labels = labels[..., 1:].contiguous()
296
+ loss = F.cross_entropy(
297
+ shift_logits.float().reshape(-1, shift_logits.size(-1)),
298
+ shift_labels.reshape(-1),
299
+ )
300
+ if self.config.z_loss_coeff:
301
+ loss = loss + self.config.z_loss_coeff * logits.float().pow(2).mean()
302
+
303
+ return CausalLMOutputWithPast(
304
+ loss=loss,
305
+ logits=logits,
306
+ past_key_values=past_key_values if use_cache else None,
307
+ )
special_tokens_map.json ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ {
2
+ "pad_token": "<|pad|>",
3
+ "bos_token": "<|bos|>",
4
+ "eos_token": "<|eos|>",
5
+ "unk_token": "<|unk|>"
6
+ }
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer_config.json ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "tokenizer_class": "PreTrainedTokenizerFast",
3
+ "model_max_length": 4096,
4
+ "pad_token": "<|pad|>",
5
+ "bos_token": "<|bos|>",
6
+ "eos_token": "<|eos|>",
7
+ "unk_token": "<|unk|>"
8
+ }