jakmro commited on
Commit
b2fa9f5
·
verified ·
1 Parent(s): e1706e4

Merge HF-format model into repo (config.json keeps JAX identity under jax_config)

Browse files
config.json CHANGED
@@ -1,5 +1,41 @@
1
  {
2
- "library_name": "jax",
3
- "model_type": "custom",
4
- "architectures": ["SimpleAttentionNetwork"]
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5
  }
 
1
  {
2
+ "architectures": [
3
+ "NeedleForCausalLM"
4
+ ],
5
+ "auto_map": {
6
+ "AutoConfig": "configuration_needle.NeedleConfig",
7
+ "AutoModel": "modeling_needle.NeedleModel",
8
+ "AutoModelForCausalLM": "modeling_needle.NeedleForCausalLM",
9
+ "AutoModelForSeq2SeqLM": "modeling_needle.NeedleForCausalLM"
10
+ },
11
+ "bos_token_id": 2,
12
+ "d_model": 512,
13
+ "decoder_start_token_id": 1,
14
+ "eos_token_id": 1,
15
+ "hidden_size": 512,
16
+ "is_encoder_decoder": true,
17
+ "model_type": "needle",
18
+ "num_attention_heads": 8,
19
+ "num_decoder_layers": 8,
20
+ "num_encoder_layers": 12,
21
+ "num_heads": 8,
22
+ "num_hidden_layers": 8,
23
+ "num_key_value_heads": 4,
24
+ "num_kv_heads": 4,
25
+ "pad_token_id": 0,
26
+ "rms_norm_eps": 1e-06,
27
+ "rope_theta": 10000.0,
28
+ "tie_word_embeddings": true,
29
+ "torch_dtype": "bfloat16",
30
+ "transformers_version": "5.5.4",
31
+ "unk_token_id": 3,
32
+ "vocab_size": 8192,
33
+ "jax_config": {
34
+ "library_name": "jax",
35
+ "model_type": "custom",
36
+ "architectures": [
37
+ "SimpleAttentionNetwork"
38
+ ],
39
+ "checkpoint": "needle.pkl"
40
+ }
41
  }
configuration_needle.py ADDED
@@ -0,0 +1,63 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Hugging Face configuration for Needle."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from transformers import PretrainedConfig
6
+
7
+
8
+ class NeedleConfig(PretrainedConfig):
9
+ model_type = "needle"
10
+
11
+ def __init__(
12
+ self,
13
+ vocab_size: int = 8192,
14
+ hidden_size: int | None = None,
15
+ d_model: int = 512,
16
+ num_attention_heads: int | None = None,
17
+ num_heads: int = 8,
18
+ num_key_value_heads: int | None = None,
19
+ num_kv_heads: int = 4,
20
+ num_encoder_layers: int = 12,
21
+ num_decoder_layers: int = 8,
22
+ rope_theta: float = 10000.0,
23
+ rms_norm_eps: float = 1e-6,
24
+ pad_token_id: int = 0,
25
+ eos_token_id: int = 1,
26
+ bos_token_id: int = 2,
27
+ unk_token_id: int = 3,
28
+ decoder_start_token_id: int | None = None,
29
+ tie_word_embeddings: bool = True,
30
+ torch_dtype: str = "bfloat16",
31
+ **kwargs,
32
+ ) -> None:
33
+ kwargs.pop("is_encoder_decoder", None)
34
+ hidden_size = int(hidden_size if hidden_size is not None else d_model)
35
+ num_attention_heads = int(num_attention_heads if num_attention_heads is not None else num_heads)
36
+ num_key_value_heads = int(num_key_value_heads if num_key_value_heads is not None else num_kv_heads)
37
+ decoder_start_token_id = eos_token_id if decoder_start_token_id is None else decoder_start_token_id
38
+
39
+ super().__init__(
40
+ pad_token_id=pad_token_id,
41
+ eos_token_id=eos_token_id,
42
+ bos_token_id=bos_token_id,
43
+ unk_token_id=unk_token_id,
44
+ decoder_start_token_id=decoder_start_token_id,
45
+ tie_word_embeddings=tie_word_embeddings,
46
+ is_encoder_decoder=True,
47
+ torch_dtype=torch_dtype,
48
+ **kwargs,
49
+ )
50
+
51
+ self.vocab_size = int(vocab_size)
52
+ self.hidden_size = hidden_size
53
+ self.d_model = hidden_size
54
+ self.num_attention_heads = num_attention_heads
55
+ self.num_heads = num_attention_heads
56
+ self.num_key_value_heads = num_key_value_heads
57
+ self.num_kv_heads = num_key_value_heads
58
+ self.num_encoder_layers = int(num_encoder_layers)
59
+ self.num_decoder_layers = int(num_decoder_layers)
60
+ self.num_hidden_layers = int(num_decoder_layers)
61
+ self.rope_theta = float(rope_theta)
62
+ self.rms_norm_eps = float(rms_norm_eps)
63
+ self.attention_head_dim = hidden_size // max(1, num_attention_heads)
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c5f9a3016e4537e492c362da5cb8ba05107d8595bec0d5ea5d8a65801db46531
3
+ size 60881792
modeling_needle.py ADDED
@@ -0,0 +1,429 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Minimal PyTorch Needle model for Cactus conversion."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+ from typing import Any
7
+
8
+ import torch
9
+ from torch import nn
10
+ import torch.nn.functional as F
11
+ from transformers import PreTrainedModel
12
+ from transformers.modeling_outputs import BaseModelOutput, Seq2SeqLMOutput
13
+
14
+ from .configuration_needle import NeedleConfig
15
+
16
+
17
+ class NeedleRMSNorm(nn.Module):
18
+ def __init__(self, hidden_size: int, eps: float) -> None:
19
+ super().__init__()
20
+ self.weight = nn.Parameter(torch.zeros(hidden_size))
21
+ self.eps = float(eps)
22
+
23
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
24
+ dtype = x.dtype
25
+ variance = x.float().pow(2).mean(dim=-1, keepdim=True)
26
+ x = x.float() * torch.rsqrt(variance + self.eps)
27
+ return x.to(dtype=dtype) * (1.0 + self.weight.to(dtype=dtype))
28
+
29
+
30
+ def _padding_mask(input_ids: torch.Tensor, pad_token_id: int) -> torch.Tensor:
31
+ return (input_ids != int(pad_token_id))[:, None, None, :]
32
+
33
+
34
+ def _causal_mask(seq_len: int, device: torch.device) -> torch.Tensor:
35
+ return torch.ones((seq_len, seq_len), dtype=torch.bool, device=device).tril()[None, None, :, :]
36
+
37
+
38
+ def _build_inv_freq(head_dim: int, theta: float) -> torch.Tensor:
39
+ return 1.0 / (float(theta) ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / float(head_dim)))
40
+
41
+
42
+ def _rotate_half(x: torch.Tensor) -> torch.Tensor:
43
+ half = x.shape[-1] // 2
44
+ return torch.cat((-x[..., half:], x[..., :half]), dim=-1)
45
+
46
+
47
+ def _apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
48
+ cos = cos.unsqueeze(2)
49
+ sin = sin.unsqueeze(2)
50
+ return (x * cos) + (_rotate_half(x) * sin)
51
+
52
+
53
+ def _rotary_tables(
54
+ inv_freq: torch.Tensor,
55
+ batch_size: int,
56
+ seq_len: int,
57
+ device: torch.device,
58
+ dtype: torch.dtype,
59
+ ) -> tuple[torch.Tensor, torch.Tensor]:
60
+ position_ids = torch.arange(seq_len, device=device, dtype=torch.long).unsqueeze(0).expand(batch_size, -1)
61
+ inv_freq = inv_freq[None, :, None].float().expand(batch_size, -1, 1).to(device)
62
+ freqs = (inv_freq @ position_ids[:, None, :].float()).transpose(1, 2)
63
+ emb = torch.cat((freqs, freqs), dim=-1)
64
+ return emb.cos().to(dtype=dtype), emb.sin().to(dtype=dtype)
65
+
66
+
67
+ def _rotary_tables_for_position_ids(
68
+ inv_freq: torch.Tensor,
69
+ position_ids: torch.Tensor,
70
+ dtype: torch.dtype,
71
+ ) -> tuple[torch.Tensor, torch.Tensor]:
72
+ inv_freq = inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1).to(position_ids.device)
73
+ freqs = (inv_freq @ position_ids[:, None, :].float()).transpose(1, 2)
74
+ emb = torch.cat((freqs, freqs), dim=-1)
75
+ return emb.cos().to(dtype=dtype), emb.sin().to(dtype=dtype)
76
+
77
+
78
+ def _add_clipped(a: torch.Tensor, b: torch.Tensor) -> torch.Tensor:
79
+ return torch.clamp(a + b, min=-65500.0, max=65500.0)
80
+
81
+
82
+ class NeedleAttention(nn.Module):
83
+ def __init__(self, config: NeedleConfig) -> None:
84
+ super().__init__()
85
+ self.hidden_size = int(config.hidden_size)
86
+ self.num_heads = int(config.num_attention_heads)
87
+ self.num_key_value_heads = int(config.num_key_value_heads)
88
+ self.head_dim = self.hidden_size // self.num_heads
89
+ kv_size = self.num_key_value_heads * self.head_dim
90
+ self.q_proj = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
91
+ self.k_proj = nn.Linear(self.hidden_size, kv_size, bias=False)
92
+ self.v_proj = nn.Linear(self.hidden_size, kv_size, bias=False)
93
+ self.out_proj = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
94
+ self.q_norm = NeedleRMSNorm(self.head_dim, config.rms_norm_eps)
95
+ self.k_norm = NeedleRMSNorm(self.head_dim, config.rms_norm_eps)
96
+ self.scale = 1.0 / math.sqrt(float(self.head_dim))
97
+
98
+ def forward(
99
+ self,
100
+ hidden_states: torch.Tensor,
101
+ key_value_states: torch.Tensor,
102
+ attention_mask: torch.Tensor | None,
103
+ rope: tuple[torch.Tensor, torch.Tensor] | None,
104
+ ) -> torch.Tensor:
105
+ batch, q_len, _ = hidden_states.shape
106
+ kv_len = key_value_states.shape[1]
107
+ q = self.q_proj(hidden_states).view(batch, q_len, self.num_heads, self.head_dim)
108
+ k = self.k_proj(key_value_states).view(batch, kv_len, self.num_key_value_heads, self.head_dim)
109
+ v = self.v_proj(key_value_states).view(batch, kv_len, self.num_key_value_heads, self.head_dim)
110
+ q = self.q_norm(q)
111
+ k = self.k_norm(k)
112
+ if rope is not None:
113
+ cos, sin = rope
114
+ q = _apply_rope(q, cos, sin)
115
+ k = _apply_rope(k, cos, sin)
116
+ out = F.scaled_dot_product_attention(
117
+ q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2),
118
+ attn_mask=attention_mask,
119
+ dropout_p=0.0,
120
+ is_causal=False,
121
+ scale=self.scale,
122
+ enable_gqa=self.num_heads != self.num_key_value_heads,
123
+ )
124
+ return self.out_proj(out.transpose(1, 2).contiguous().view(batch, q_len, self.hidden_size))
125
+
126
+ def project_kv(self, key_value_states: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
127
+ batch, kv_len, _ = key_value_states.shape
128
+ k = self.k_proj(key_value_states).view(batch, kv_len, self.num_key_value_heads, self.head_dim)
129
+ v = self.v_proj(key_value_states).view(batch, kv_len, self.num_key_value_heads, self.head_dim)
130
+ k = self.k_norm(k)
131
+ return k.contiguous(), v.contiguous()
132
+
133
+ def forward_with_kv(
134
+ self,
135
+ hidden_states: torch.Tensor,
136
+ key_states: torch.Tensor,
137
+ value_states: torch.Tensor,
138
+ attention_mask: torch.Tensor | None,
139
+ rope: tuple[torch.Tensor, torch.Tensor] | None,
140
+ ) -> torch.Tensor:
141
+ batch, q_len, _ = hidden_states.shape
142
+ q = self.q_proj(hidden_states).view(batch, q_len, self.num_heads, self.head_dim)
143
+ q = self.q_norm(q)
144
+ if rope is not None:
145
+ cos, sin = rope
146
+ q = _apply_rope(q, cos, sin)
147
+ out = F.scaled_dot_product_attention(
148
+ q.transpose(1, 2), key_states.transpose(1, 2), value_states.transpose(1, 2),
149
+ attn_mask=attention_mask,
150
+ dropout_p=0.0,
151
+ is_causal=False,
152
+ scale=self.scale,
153
+ enable_gqa=self.num_heads != self.num_key_value_heads,
154
+ )
155
+ return self.out_proj(out.transpose(1, 2).contiguous().view(batch, q_len, self.hidden_size))
156
+
157
+
158
+ class NeedleEncoderLayer(nn.Module):
159
+ def __init__(self, config: NeedleConfig) -> None:
160
+ super().__init__()
161
+ self.input_layernorm = NeedleRMSNorm(config.hidden_size, config.rms_norm_eps)
162
+ self.self_attn = NeedleAttention(config)
163
+ self.attn_gate = nn.Parameter(torch.zeros(1))
164
+
165
+ def forward(
166
+ self,
167
+ hidden_states: torch.Tensor,
168
+ attention_mask: torch.Tensor,
169
+ rope: tuple[torch.Tensor, torch.Tensor],
170
+ ) -> torch.Tensor:
171
+ normed = self.input_layernorm(hidden_states)
172
+ attn = self.self_attn(normed, normed, attention_mask, rope)
173
+ return _add_clipped(hidden_states, torch.sigmoid(self.attn_gate).to(dtype=attn.dtype) * attn)
174
+
175
+
176
+ class NeedleDecoderLayer(nn.Module):
177
+ def __init__(self, config: NeedleConfig) -> None:
178
+ super().__init__()
179
+ self.input_layernorm = NeedleRMSNorm(config.hidden_size, config.rms_norm_eps)
180
+ self.self_attn = NeedleAttention(config)
181
+ self.self_attn_gate = nn.Parameter(torch.zeros(1))
182
+ self.encoder_attn_layer_norm = NeedleRMSNorm(config.hidden_size, config.rms_norm_eps)
183
+ self.encoder_attn = NeedleAttention(config)
184
+ self.cross_attn_gate = nn.Parameter(torch.zeros(1))
185
+
186
+ def forward(
187
+ self,
188
+ hidden_states: torch.Tensor,
189
+ encoder_hidden_states: torch.Tensor,
190
+ self_mask: torch.Tensor,
191
+ encoder_mask: torch.Tensor,
192
+ rope: tuple[torch.Tensor, torch.Tensor],
193
+ ) -> torch.Tensor:
194
+ normed = self.input_layernorm(hidden_states)
195
+ attn = self.self_attn(normed, normed, self_mask, rope)
196
+ hidden_states = _add_clipped(hidden_states, torch.sigmoid(self.self_attn_gate).to(dtype=attn.dtype) * attn)
197
+ attn = self.encoder_attn(self.encoder_attn_layer_norm(hidden_states), encoder_hidden_states, encoder_mask, None)
198
+ return _add_clipped(hidden_states, torch.sigmoid(self.cross_attn_gate).to(dtype=attn.dtype) * attn)
199
+
200
+
201
+ class NeedleEncoder(nn.Module):
202
+ def __init__(self, config: NeedleConfig) -> None:
203
+ super().__init__()
204
+ self.layers = nn.ModuleList([NeedleEncoderLayer(config) for _ in range(config.num_encoder_layers)])
205
+ self.final_norm = NeedleRMSNorm(config.hidden_size, config.rms_norm_eps)
206
+ self.head_dim = config.hidden_size // config.num_attention_heads
207
+ self.rope_theta = float(config.rope_theta)
208
+ self.register_buffer("inv_freq", _build_inv_freq(self.head_dim, self.rope_theta), persistent=False)
209
+
210
+ def reset_rope(self) -> None:
211
+ self.inv_freq = _build_inv_freq(self.head_dim, self.rope_theta).to(device=self.inv_freq.device)
212
+
213
+ def forward(self, hidden_states: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor:
214
+ inv_freq = _build_inv_freq(self.head_dim, self.rope_theta).to(device=hidden_states.device)
215
+ rope = _rotary_tables(
216
+ inv_freq,
217
+ hidden_states.shape[0],
218
+ hidden_states.shape[1],
219
+ hidden_states.device,
220
+ hidden_states.dtype,
221
+ )
222
+ for layer in self.layers:
223
+ hidden_states = layer(hidden_states, attention_mask, rope)
224
+ return self.final_norm(hidden_states)
225
+
226
+
227
+ class NeedleDecoder(nn.Module):
228
+ def __init__(self, config: NeedleConfig) -> None:
229
+ super().__init__()
230
+ self.layers = nn.ModuleList([NeedleDecoderLayer(config) for _ in range(config.num_decoder_layers)])
231
+ self.norm = NeedleRMSNorm(config.hidden_size, config.rms_norm_eps)
232
+ self.head_dim = config.hidden_size // config.num_attention_heads
233
+ self.rope_theta = float(config.rope_theta)
234
+ self.register_buffer("inv_freq", _build_inv_freq(self.head_dim, self.rope_theta), persistent=False)
235
+
236
+ def reset_rope(self) -> None:
237
+ self.inv_freq = _build_inv_freq(self.head_dim, self.rope_theta).to(device=self.inv_freq.device)
238
+
239
+ def forward(
240
+ self,
241
+ hidden_states: torch.Tensor,
242
+ encoder_hidden_states: torch.Tensor,
243
+ self_mask: torch.Tensor,
244
+ encoder_mask: torch.Tensor,
245
+ ) -> torch.Tensor:
246
+ inv_freq = _build_inv_freq(self.head_dim, self.rope_theta).to(device=hidden_states.device)
247
+ rope = _rotary_tables(
248
+ inv_freq,
249
+ hidden_states.shape[0],
250
+ hidden_states.shape[1],
251
+ hidden_states.device,
252
+ hidden_states.dtype,
253
+ )
254
+ for layer in self.layers:
255
+ hidden_states = layer(hidden_states, encoder_hidden_states, self_mask, encoder_mask, rope)
256
+ return self.norm(hidden_states)
257
+
258
+
259
+ class NeedleModel(PreTrainedModel):
260
+ config_class = NeedleConfig
261
+ base_model_prefix = "model"
262
+ main_input_name = "input_ids"
263
+
264
+ def __init__(self, config: NeedleConfig) -> None:
265
+ super().__init__(config)
266
+ self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size)
267
+ self.embed_scale = math.sqrt(float(config.hidden_size))
268
+ self.encoder = NeedleEncoder(config)
269
+ self.decoder = NeedleDecoder(config)
270
+ self.post_init()
271
+ self.reset_rope()
272
+
273
+ def reset_rope(self) -> None:
274
+ self.encoder.reset_rope()
275
+ self.decoder.reset_rope()
276
+
277
+ def get_input_embeddings(self) -> nn.Embedding:
278
+ return self.embed_tokens
279
+
280
+ def set_input_embeddings(self, value: nn.Embedding) -> None:
281
+ self.embed_tokens = value
282
+
283
+ def forward(
284
+ self,
285
+ input_ids: torch.Tensor,
286
+ attention_mask: torch.Tensor | None = None,
287
+ decoder_input_ids: torch.Tensor | None = None,
288
+ **_: Any,
289
+ ) -> BaseModelOutput:
290
+ decoder_input_ids = input_ids if decoder_input_ids is None else decoder_input_ids
291
+ encoder_mask = _padding_mask(input_ids, self.config.pad_token_id)
292
+ if attention_mask is not None:
293
+ encoder_mask = encoder_mask & attention_mask[:, None, None, :].to(dtype=torch.bool)
294
+ encoder_hidden = self.embed_tokens(input_ids) * self.embed_scale
295
+ encoder_hidden = self.encoder(encoder_hidden, encoder_mask)
296
+
297
+ self_mask = _causal_mask(decoder_input_ids.shape[1], decoder_input_ids.device)
298
+ decoder_hidden = self.embed_tokens(decoder_input_ids) * self.embed_scale
299
+ decoder_hidden = self.decoder(decoder_hidden, encoder_hidden, self_mask, encoder_mask)
300
+ return BaseModelOutput(last_hidden_state=decoder_hidden)
301
+
302
+
303
+ class NeedleForCausalLM(PreTrainedModel):
304
+ config_class = NeedleConfig
305
+ base_model_prefix = "model"
306
+ main_input_name = "input_ids"
307
+ _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"}
308
+
309
+ def __init__(self, config: NeedleConfig) -> None:
310
+ super().__init__(config)
311
+ self.model = NeedleModel(config)
312
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
313
+ self.post_init()
314
+ self.model.reset_rope()
315
+ self.tie_weights()
316
+
317
+ def get_encoder(self) -> NeedleEncoder:
318
+ return self.model.encoder
319
+
320
+ def get_input_embeddings(self) -> nn.Embedding:
321
+ return self.model.embed_tokens
322
+
323
+ def set_input_embeddings(self, value: nn.Embedding) -> None:
324
+ self.model.embed_tokens = value
325
+
326
+ def get_output_embeddings(self) -> nn.Linear:
327
+ return self.lm_head
328
+
329
+ def set_output_embeddings(self, value: nn.Linear) -> None:
330
+ self.lm_head = value
331
+
332
+ def tie_weights(self, *args: Any, **kwargs: Any) -> None:
333
+ del args, kwargs
334
+ if self.config.tie_word_embeddings:
335
+ self.lm_head.weight = self.model.embed_tokens.weight
336
+
337
+ def forward(
338
+ self,
339
+ input_ids: torch.Tensor,
340
+ attention_mask: torch.Tensor | None = None,
341
+ decoder_input_ids: torch.Tensor | None = None,
342
+ **kwargs: Any,
343
+ ) -> Seq2SeqLMOutput:
344
+ hidden_states = self.model(
345
+ input_ids=input_ids,
346
+ attention_mask=attention_mask,
347
+ decoder_input_ids=decoder_input_ids,
348
+ **kwargs,
349
+ ).last_hidden_state
350
+ return Seq2SeqLMOutput(logits=self.lm_head(hidden_states))
351
+
352
+ def cactus_source_encode(
353
+ self,
354
+ input_ids: torch.Tensor,
355
+ attention_mask: torch.Tensor,
356
+ ) -> tuple[torch.Tensor, torch.Tensor]:
357
+ base_mask = _padding_mask(input_ids, self.config.pad_token_id)
358
+ base_mask = base_mask & attention_mask[:, None, None, :].to(dtype=torch.bool)
359
+ encoder_mask = base_mask.expand(
360
+ -1,
361
+ int(self.config.num_attention_heads),
362
+ input_ids.shape[1],
363
+ -1,
364
+ ).contiguous()
365
+ encoder_hidden = self.model.embed_tokens(input_ids) * self.model.embed_scale
366
+ encoder_hidden = self.model.encoder(encoder_hidden, encoder_mask)
367
+ decoder_mask = base_mask.expand(
368
+ -1,
369
+ int(self.config.num_attention_heads),
370
+ 1,
371
+ -1,
372
+ ).contiguous()
373
+ return encoder_hidden, decoder_mask.to(dtype=encoder_hidden.dtype)
374
+
375
+ def cactus_decoder_cross_kv(
376
+ self,
377
+ encoder_hidden_states: torch.Tensor,
378
+ encoder_attention_mask: torch.Tensor,
379
+ ) -> tuple[torch.Tensor, ...]:
380
+ del encoder_attention_mask
381
+ outputs: list[torch.Tensor] = []
382
+ for layer in self.model.decoder.layers:
383
+ k, v = layer.encoder_attn.project_kv(encoder_hidden_states)
384
+ outputs.extend((k, v))
385
+ return tuple(outputs)
386
+
387
+ def cactus_decoder_step(
388
+ self,
389
+ decoder_input_ids: torch.Tensor,
390
+ position_ids: torch.Tensor,
391
+ encoder_attention_mask: torch.Tensor,
392
+ *cross_kv: torch.Tensor,
393
+ ) -> torch.Tensor:
394
+ encoder_attention_mask = encoder_attention_mask != 0
395
+ hidden_states = self.model.embed_tokens(decoder_input_ids) * self.model.embed_scale
396
+ inv_freq = _build_inv_freq(
397
+ self.model.decoder.head_dim,
398
+ self.model.decoder.rope_theta,
399
+ ).to(device=hidden_states.device)
400
+ rope = _rotary_tables_for_position_ids(
401
+ inv_freq,
402
+ position_ids.to(dtype=torch.long),
403
+ hidden_states.dtype,
404
+ )
405
+ for layer_index, layer in enumerate(self.model.decoder.layers):
406
+ normed = layer.input_layernorm(hidden_states)
407
+ attn = layer.self_attn(normed, normed, None, rope)
408
+ hidden_states = _add_clipped(hidden_states, torch.sigmoid(layer.self_attn_gate).to(dtype=attn.dtype) * attn)
409
+
410
+ cross_attn = layer.encoder_attn.forward_with_kv(
411
+ layer.encoder_attn_layer_norm(hidden_states),
412
+ cross_kv[layer_index * 2],
413
+ cross_kv[layer_index * 2 + 1],
414
+ encoder_attention_mask,
415
+ None,
416
+ )
417
+ hidden_states = _add_clipped(hidden_states, torch.sigmoid(layer.cross_attn_gate).to(dtype=cross_attn.dtype) * cross_attn)
418
+ hidden_states = self.model.decoder.norm(hidden_states)
419
+ return self.lm_head(hidden_states)
420
+
421
+ def _init_weights(self, module: nn.Module) -> None:
422
+ if isinstance(module, nn.Linear):
423
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
424
+ if module.bias is not None:
425
+ nn.init.zeros_(module.bias)
426
+ elif isinstance(module, nn.Embedding):
427
+ nn.init.normal_(module.weight, mean=0.0, std=0.02)
428
+ elif isinstance(module, NeedleRMSNorm):
429
+ nn.init.zeros_(module.weight)
special_tokens_map.json ADDED
@@ -0,0 +1,11 @@
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "additional_special_tokens": [
3
+ "<tool_call>",
4
+ "<tools>"
5
+ ],
6
+ "bos_token": "<s>",
7
+ "eos_token": "</s>",
8
+ "pad_token": "<pad>",
9
+ "unk_token": "<unk>"
10
+ }
11
+
tokenization_needle.py ADDED
@@ -0,0 +1,121 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Slow SentencePiece tokenizer for Needle."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import os
6
+ import shutil
7
+ from typing import Any
8
+
9
+ import sentencepiece as spm
10
+ from transformers import PreTrainedTokenizer
11
+
12
+
13
+ VOCAB_FILES_NAMES = {"vocab_file": "tokenizer.model"}
14
+
15
+
16
+ class NeedleTokenizer(PreTrainedTokenizer):
17
+ vocab_files_names = VOCAB_FILES_NAMES
18
+ model_input_names = ["input_ids", "attention_mask"]
19
+
20
+ def __init__(
21
+ self,
22
+ vocab_file: str,
23
+ unk_token: str = "<unk>",
24
+ bos_token: str = "<s>",
25
+ eos_token: str = "</s>",
26
+ pad_token: str = "<pad>",
27
+ tool_call_token: str = "<tool_call>",
28
+ tools_token: str = "<tools>",
29
+ **kwargs: Any,
30
+ ) -> None:
31
+ self.vocab_file = vocab_file
32
+ self.sp_model = spm.SentencePieceProcessor()
33
+ self.sp_model.Load(vocab_file)
34
+ self.sp = self.sp_model
35
+ self.tool_call_token = tool_call_token
36
+ self.tools_token = tools_token
37
+ additional = list(kwargs.pop("additional_special_tokens", []) or [])
38
+ for token in (tool_call_token, tools_token):
39
+ if token not in additional:
40
+ additional.append(token)
41
+ super().__init__(
42
+ unk_token=unk_token,
43
+ bos_token=bos_token,
44
+ eos_token=eos_token,
45
+ pad_token=pad_token,
46
+ additional_special_tokens=additional,
47
+ **kwargs,
48
+ )
49
+
50
+ @property
51
+ def vocab_size(self) -> int:
52
+ return int(self.sp_model.GetPieceSize())
53
+
54
+ @property
55
+ def tool_call_token_id(self) -> int:
56
+ return int(self.sp_model.PieceToId(self.tool_call_token))
57
+
58
+ @property
59
+ def tools_token_id(self) -> int:
60
+ return int(self.sp_model.PieceToId(self.tools_token))
61
+
62
+ def get_vocab(self) -> dict[str, int]:
63
+ vocab = {self.sp_model.IdToPiece(i): i for i in range(self.vocab_size)}
64
+ vocab.update(self.added_tokens_encoder)
65
+ return vocab
66
+
67
+ def _tokenize(self, text: str) -> list[str]:
68
+ return list(self.sp_model.EncodeAsPieces(text))
69
+
70
+ def _convert_token_to_id(self, token: str) -> int:
71
+ return int(self.sp_model.PieceToId(token))
72
+
73
+ def _convert_id_to_token(self, index: int) -> str:
74
+ return str(self.sp_model.IdToPiece(int(index)))
75
+
76
+ def convert_tokens_to_string(self, tokens: list[str]) -> str:
77
+ return self.sp_model.DecodePieces(tokens)
78
+
79
+ def build_inputs_with_special_tokens(
80
+ self,
81
+ token_ids_0: list[int],
82
+ token_ids_1: list[int] | None = None,
83
+ ) -> list[int]:
84
+ if token_ids_1 is None:
85
+ return list(token_ids_0)
86
+ return list(token_ids_0) + list(token_ids_1)
87
+
88
+ def get_special_tokens_mask(
89
+ self,
90
+ token_ids_0: list[int],
91
+ token_ids_1: list[int] | None = None,
92
+ already_has_special_tokens: bool = False,
93
+ ) -> list[int]:
94
+ if already_has_special_tokens:
95
+ all_ids = list(token_ids_0)
96
+ else:
97
+ all_ids = self.build_inputs_with_special_tokens(token_ids_0, token_ids_1)
98
+ special = {
99
+ self.pad_token_id,
100
+ self.eos_token_id,
101
+ self.bos_token_id,
102
+ self.unk_token_id,
103
+ self.tool_call_token_id,
104
+ self.tools_token_id,
105
+ }
106
+ return [1 if token_id in special else 0 for token_id in all_ids]
107
+
108
+ def create_token_type_ids_from_sequences(
109
+ self,
110
+ token_ids_0: list[int],
111
+ token_ids_1: list[int] | None = None,
112
+ ) -> list[int]:
113
+ return [0] * len(self.build_inputs_with_special_tokens(token_ids_0, token_ids_1))
114
+
115
+ def save_vocabulary(self, save_directory: str, filename_prefix: str | None = None) -> tuple[str]:
116
+ os.makedirs(save_directory, exist_ok=True)
117
+ out_name = "tokenizer.model" if filename_prefix is None else f"{filename_prefix}-tokenizer.model"
118
+ out_path = os.path.join(save_directory, out_name)
119
+ if os.path.abspath(self.vocab_file) != os.path.abspath(out_path):
120
+ shutil.copyfile(self.vocab_file, out_path)
121
+ return (out_path,)
tokenizer.model ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0823f5b9133c68a8140addc5d7a425fa9119c4c8cb4a550363b4bffa4ba1c8c7
3
+ size 124960
tokenizer_config.json ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "additional_special_tokens": [
3
+ "<tool_call>",
4
+ "<tools>"
5
+ ],
6
+ "auto_map": {
7
+ "AutoTokenizer": [
8
+ "tokenization_needle.NeedleTokenizer",
9
+ null
10
+ ]
11
+ },
12
+ "bos_token": "<s>",
13
+ "clean_up_tokenization_spaces": false,
14
+ "eos_token": "</s>",
15
+ "model_input_names": [
16
+ "input_ids",
17
+ "attention_mask"
18
+ ],
19
+ "model_max_length": 1024,
20
+ "pad_token": "<pad>",
21
+ "padding_side": "right",
22
+ "tokenizer_class": "NeedleTokenizer",
23
+ "tool_call_token": "<tool_call>",
24
+ "tools_token": "<tools>",
25
+ "unk_token": "<unk>"
26
+ }
27
+