xlr8harder commited on
Commit
c4d4d45
·
verified ·
1 Parent(s): bfb5c3c

Upload Talkie YaRN 32k step500 checkpoint

Browse files
README.md ADDED
@@ -0,0 +1,118 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ language:
3
+ - en
4
+ license: apache-2.0
5
+ library_name: transformers
6
+ pipeline_tag: text-generation
7
+ base_model: xlr8harder/talkie-1930-13b-base-tf
8
+ datasets:
9
+ - common-pile/project_gutenberg_filtered
10
+ tags:
11
+ - transformers
12
+ - safetensors
13
+ - bfloat16
14
+ - custom_code
15
+ - text-generation
16
+ - talkie
17
+ - yarn
18
+ - long-context
19
+ - pre-1931
20
+ ---
21
+
22
+ # Talkie 1930 13B YaRN 32k
23
+
24
+ This is a 32k-context YaRN extension of
25
+ [`xlr8harder/talkie-1930-13b-base-tf`](https://huggingface.co/xlr8harder/talkie-1930-13b-base-tf).
26
+ It is the recommended long-context checkpoint from this experiment series.
27
+
28
+ The checkpoint was made by applying YaRN with a 16x extension from the 2,048-token
29
+ configuration in the reference Talkie repository, then continuing pretraining for
30
+ 500 steps at 32,768 tokens. The continued pretraining data was a Project Gutenberg
31
+ split filtered to English public-domain books with publication years 1500-1930,
32
+ for 265,080,702 Talkie tokens. Training used 262,144 tokens per step, cosine LR
33
+ decay from `1e-5` to `1e-6`, 50 warmup steps, and weight decay `0.01`.
34
+
35
+ We originally used a 2k starting context because the public reference config
36
+ advertised 2,048 positions. The Talkie team later clarified that the base model
37
+ had been trained with 4k context. We also ran a 4k-start, 8x-extension variant;
38
+ it was slightly stronger at short contexts but substantially weaker at 32k and
39
+ collapsed on RULER variable tracking. That alternate checkpoint is published as
40
+ [`xlr8harder/talkie-1930-13b-yarn-32k-from4k-step1000-tf`](https://huggingface.co/xlr8harder/talkie-1930-13b-yarn-32k-from4k-step1000-tf).
41
+
42
+ We selected step500 because it was more well rounded than the final step1000
43
+ checkpoint from the same 2k-start run.
44
+
45
+ ## Usage
46
+
47
+ This model uses custom Talkie modeling/tokenization code, so load it with
48
+ `trust_remote_code=True`.
49
+
50
+ ```python
51
+ from transformers import AutoModelForCausalLM, AutoTokenizer
52
+
53
+ model_id = "xlr8harder/talkie-1930-13b-yarn-32k-tf"
54
+
55
+ tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
56
+ model = AutoModelForCausalLM.from_pretrained(
57
+ model_id,
58
+ torch_dtype="auto",
59
+ device_map="auto",
60
+ trust_remote_code=True,
61
+ )
62
+ ```
63
+
64
+ For vLLM, set `--max-model-len 32768` and enable remote code.
65
+
66
+ ## RULER Results
67
+
68
+ Scores are aggregate RULER accuracy percentages from our harness, using 100
69
+ examples per task, greedy decoding, and the same prompt generation setup within
70
+ each model family. Different tokenizers mean nominal context lengths are not
71
+ byte-identical across unrelated model families, so use open-model rows as
72
+ orientation rather than exact head-to-head leaderboard claims.
73
+
74
+ | Model / setup | 2k | 4k | 8k | 16k | 32k |
75
+ | --- | ---: | ---: | ---: | ---: | ---: |
76
+ | Talkie base, extrapolation only | 85.86 | 77.71 | 23.40 | n/a | n/a |
77
+ | Talkie YaRN 32k, 2k start, step500 | 80.78 | 79.50 | 73.15 | 70.05 | 61.83 |
78
+ | Talkie YaRN 32k, 2k start, step1000 | 80.30 | 79.94 | 73.17 | 67.98 | 61.83 |
79
+ | Talkie YaRN 32k, 4k start, step500 | 83.80 | 80.71 | 75.64 | 68.80 | 54.76 |
80
+ | Talkie YaRN 32k, 4k start, step1000 | 84.18 | 80.98 | 76.17 | 68.45 | 55.01 |
81
+ | Llama 3.1 8B base | 97.12 | 94.25 | 92.34 | 91.61 | 88.54 |
82
+ | Yarn-Llama-2 13B 64k | 90.78 | 81.95 | 70.39 | 60.02 | 52.60 |
83
+ | Qwen3 8B pretrain base | 98.90 | 95.83 | 94.37 | 93.04 | 89.39 |
84
+
85
+ At 32k, the 2k-start step500 checkpoint was meaningfully stronger than the
86
+ 4k-start checkpoints despite the 4k-start checkpoints being better at shorter
87
+ lengths. The largest qualitative difference was variable tracking (`vt`), where
88
+ the 4k-start run collapsed to near zero while this checkpoint retained partial
89
+ ability.
90
+
91
+ ## 32k Per-Task RULER Breakdown
92
+
93
+ | Task | Score |
94
+ | --- | ---: |
95
+ | `cwe` | 15.90 |
96
+ | `fwe` | 34.67 |
97
+ | `niah_multikey_1` | 97.00 |
98
+ | `niah_multikey_2` | 98.00 |
99
+ | `niah_multikey_3` | 16.00 |
100
+ | `niah_multiquery` | 92.25 |
101
+ | `niah_multivalue` | 54.75 |
102
+ | `niah_single_1` | 100.00 |
103
+ | `niah_single_2` | 100.00 |
104
+ | `niah_single_3` | 78.00 |
105
+ | `qa_1` | 49.00 |
106
+ | `qa_2` | 42.00 |
107
+ | `vt` | 26.20 |
108
+
109
+ Task shorthand: `vt` is variable tracking, `cwe` is common-word extraction,
110
+ `fwe` is frequent/coded-word extraction, `niah_*` are needle-in-a-haystack
111
+ retrieval variants, and `qa_*` are long-context question-answering tasks.
112
+
113
+ ## Notes
114
+
115
+ This is a research checkpoint for long-context experimentation. It improves
116
+ Talkie's long-context RULER behavior relative to pure extrapolation, but it does
117
+ not match stronger modern long-context baselines. Use normal evaluation for your
118
+ target workload before relying on 32k behavior.
checkpoint_complete.json ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ {
2
+ "checkpoint": "model-step-000500",
3
+ "checkpoint_type": "model",
4
+ "dtype": "bfloat16",
5
+ "format": "transformers-safetensors",
6
+ "time_unix": 1778770658.2032046
7
+ }
config.json ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "TalkieForCausalLM"
4
+ ],
5
+ "auto_map": {
6
+ "AutoConfig": "configuration_talkie.TalkieConfig",
7
+ "AutoModel": "modeling_talkie.TalkieModel",
8
+ "AutoModelForCausalLM": "modeling_talkie.TalkieForCausalLM"
9
+ },
10
+ "bos_token_id": null,
11
+ "dtype": "bfloat16",
12
+ "eos_token_id": 65535,
13
+ "head_dim": 128,
14
+ "hidden_size": 5120,
15
+ "logit_scale": 1.0,
16
+ "max_position_embeddings": 32768,
17
+ "model_type": "talkie",
18
+ "n_embd": 5120,
19
+ "n_head": 40,
20
+ "n_layer": 40,
21
+ "num_attention_heads": 40,
22
+ "num_hidden_layers": 40,
23
+ "pad_token_id": 65535,
24
+ "rope_base": 1000000,
25
+ "rope_parameters": {
26
+ "beta_fast": 32.0,
27
+ "beta_slow": 1.0,
28
+ "factor": 16.0,
29
+ "original_max_position_embeddings": 2048,
30
+ "rope_type": "yarn"
31
+ },
32
+ "style": "base",
33
+ "tie_word_embeddings": false,
34
+ "transformers_version": "5.8.1",
35
+ "use_cache": true,
36
+ "vocab_size": 65536,
37
+ "torch_dtype": "bfloat16"
38
+ }
configuration_talkie.py ADDED
@@ -0,0 +1,100 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ from collections.abc import Mapping
4
+
5
+ from transformers import PretrainedConfig
6
+
7
+
8
+ class TalkieConfig(PretrainedConfig):
9
+ model_type = "talkie"
10
+
11
+ def __init__(
12
+ self,
13
+ vocab_size: int = 65536,
14
+ n_layer: int = 40,
15
+ n_head: int = 40,
16
+ n_embd: int = 5120,
17
+ head_dim: int = 128,
18
+ max_position_embeddings: int = 2048,
19
+ rope_base: int = 1_000_000,
20
+ rope_scaling: dict | None = None,
21
+ rope_parameters: dict | None = None,
22
+ logit_scale: float = 1.0,
23
+ use_cache: bool = True,
24
+ tie_word_embeddings: bool = False,
25
+ bos_token_id: int | None = None,
26
+ eos_token_id: int | list[int] = 65535,
27
+ pad_token_id: int | None = None,
28
+ **kwargs,
29
+ ):
30
+ if rope_scaling is None:
31
+ rope_scaling = rope_parameters
32
+ self.max_position_embeddings = max_position_embeddings
33
+ self.rope_scaling = self._normalize_rope_scaling(rope_scaling)
34
+ self.rope_parameters = self.rope_scaling
35
+ super().__init__(
36
+ bos_token_id=bos_token_id,
37
+ eos_token_id=eos_token_id,
38
+ pad_token_id=pad_token_id,
39
+ tie_word_embeddings=tie_word_embeddings,
40
+ **kwargs,
41
+ )
42
+ self.vocab_size = vocab_size
43
+ self.n_layer = n_layer
44
+ self.n_head = n_head
45
+ self.n_embd = n_embd
46
+ self.head_dim = head_dim
47
+ self.max_position_embeddings = max_position_embeddings
48
+ self.rope_base = rope_base
49
+ self.rope_scaling = self._normalize_rope_scaling(rope_scaling)
50
+ self.rope_parameters = self.rope_scaling
51
+ self.logit_scale = logit_scale
52
+ self.use_cache = use_cache
53
+
54
+ # Common Transformers aliases used by generation/cache helpers.
55
+ self.hidden_size = n_embd
56
+ self.num_hidden_layers = n_layer
57
+ self.num_attention_heads = n_head
58
+
59
+ @staticmethod
60
+ def _normalize_rope_scaling(rope_scaling: dict | None) -> dict | None:
61
+ if rope_scaling is None:
62
+ return None
63
+ if not isinstance(rope_scaling, Mapping):
64
+ raise TypeError("rope_scaling must be a dictionary")
65
+
66
+ scaling = dict(rope_scaling)
67
+ rope_type = scaling.get("rope_type", scaling.get("type"))
68
+ if rope_type is None:
69
+ raise ValueError("rope_scaling must include 'rope_type' or 'type'")
70
+
71
+ rope_type = str(rope_type).lower()
72
+ if rope_type == "ntk":
73
+ rope_type = "dynamic"
74
+ supported = {"default", "linear", "dynamic", "yarn"}
75
+ if rope_type not in supported:
76
+ raise ValueError(
77
+ f"unsupported rope_scaling type {rope_type!r}; expected one of {sorted(supported)}"
78
+ )
79
+
80
+ if rope_type == "default":
81
+ return None
82
+
83
+ factor = float(scaling.get("factor", 1.0))
84
+ if factor < 1.0:
85
+ raise ValueError("rope_scaling factor must be >= 1.0")
86
+
87
+ scaling["rope_type"] = rope_type
88
+ scaling.pop("type", None)
89
+ scaling["factor"] = factor
90
+ if "original_max_position_embeddings" in scaling:
91
+ scaling["original_max_position_embeddings"] = int(
92
+ scaling["original_max_position_embeddings"]
93
+ )
94
+ if "beta_fast" in scaling:
95
+ scaling["beta_fast"] = float(scaling["beta_fast"])
96
+ if "beta_slow" in scaling:
97
+ scaling["beta_slow"] = float(scaling["beta_slow"])
98
+ if "attention_factor" in scaling and scaling["attention_factor"] is not None:
99
+ scaling["attention_factor"] = float(scaling["attention_factor"])
100
+ return scaling
generation_config.json ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "do_sample": true,
3
+ "eos_token_id": 65535,
4
+ "pad_token_id": 65535,
5
+ "temperature": 0.7,
6
+ "transformers_version": "5.8.1",
7
+ "use_cache": true
8
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:36c525ac4b08f2aadd7fe7b12ecf8b4bcb1dbceb9e08e5e57b2256c586550aee
3
+ size 26560480408
modeling_talkie.py ADDED
@@ -0,0 +1,630 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import math
4
+
5
+ import torch
6
+ import torch.nn as nn
7
+ import torch.nn.functional as F
8
+ from transformers.cache_utils import Cache, DynamicCache
9
+ from transformers import GenerationMixin
10
+ from transformers import PreTrainedModel
11
+ from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast
12
+
13
+ try:
14
+ from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS
15
+ except ImportError: # pragma: no cover - compatibility with older Transformers.
16
+ ALL_ATTENTION_FUNCTIONS = None
17
+
18
+ from .configuration_talkie import TalkieConfig
19
+
20
+
21
+ def eager_attention_forward(
22
+ module: nn.Module,
23
+ query: torch.Tensor,
24
+ key: torch.Tensor,
25
+ value: torch.Tensor,
26
+ attention_mask: torch.Tensor | None,
27
+ dropout: float = 0.0,
28
+ scaling: float | None = None,
29
+ is_causal: bool | None = None,
30
+ **kwargs,
31
+ ) -> tuple[torch.Tensor, None]:
32
+ del kwargs
33
+ is_causal = is_causal if is_causal is not None else getattr(module, "is_causal", True)
34
+ output = F.scaled_dot_product_attention(
35
+ query,
36
+ key,
37
+ value,
38
+ attn_mask=attention_mask,
39
+ dropout_p=dropout,
40
+ scale=scaling,
41
+ is_causal=is_causal and attention_mask is None,
42
+ )
43
+ return output.transpose(1, 2).contiguous(), None
44
+
45
+
46
+ def apply_rotary_emb(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
47
+ d = x.shape[3] // 2
48
+ x1 = x[..., :d]
49
+ x2 = x[..., d:]
50
+ y1 = x1 * cos + x2 * sin
51
+ y2 = x1 * (-sin) + x2 * cos
52
+ return torch.cat([y1, y2], 3).type_as(x)
53
+
54
+
55
+ class HeadGain(nn.Module):
56
+ def __init__(self, n_head: int):
57
+ super().__init__()
58
+ self.head_g = nn.Parameter(torch.ones([n_head]))
59
+
60
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
61
+ return x * self.head_g.type_as(x).view(1, 1, -1, 1)
62
+
63
+
64
+ class WeightGain(nn.Module):
65
+ def __init__(self):
66
+ super().__init__()
67
+ self.w_g = nn.Parameter(torch.ones(1))
68
+
69
+ def forward(self, w: torch.Tensor) -> torch.Tensor:
70
+ return w * self.w_g.type_as(w)
71
+
72
+
73
+ class ActGain(nn.Module):
74
+ def __init__(self, init_value: float):
75
+ super().__init__()
76
+ self.a_g = nn.Parameter(torch.ones(1) * init_value)
77
+
78
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
79
+ return x * self.a_g.type_as(x)
80
+
81
+
82
+ class CausalSelfAttention(nn.Module):
83
+ is_causal = True
84
+
85
+ def __init__(self, config: TalkieConfig, layer_idx: int):
86
+ super().__init__()
87
+ self.config = config
88
+ self.layer_idx = layer_idx
89
+ self.n_head = config.n_head
90
+ self.head_dim = config.head_dim
91
+ n_state = config.n_embd
92
+
93
+ self.attn_query = nn.Linear(n_state, n_state, bias=False)
94
+ self.attn_key = nn.Linear(n_state, n_state, bias=False)
95
+ self.attn_value = nn.Linear(n_state, n_state, bias=False)
96
+ self.attn_resid = nn.Linear(n_state, n_state, bias=False)
97
+ self.head_gain = HeadGain(config.n_head)
98
+
99
+ def forward(
100
+ self,
101
+ x: torch.Tensor,
102
+ cos_sin: tuple[torch.Tensor, torch.Tensor],
103
+ attention_mask: torch.Tensor | None = None,
104
+ **kwargs,
105
+ ) -> torch.Tensor:
106
+ bsz, seq_len, _ = x.size()
107
+ q = self.attn_query(x).view(bsz, seq_len, self.n_head, self.head_dim)
108
+ k = self.attn_key(x).view(bsz, seq_len, self.n_head, self.head_dim)
109
+ v = self.attn_value(x).view(bsz, seq_len, self.n_head, self.head_dim)
110
+
111
+ cos, sin = cos_sin
112
+ q, k = apply_rotary_emb(q, cos, sin), apply_rotary_emb(k, cos, sin)
113
+ q, k = F.rms_norm(q, (q.size(-1),)), F.rms_norm(k, (k.size(-1),))
114
+ q = self.head_gain(q)
115
+
116
+ key_states = k.transpose(1, 2)
117
+ value_states = v.transpose(1, 2)
118
+ if kwargs.get("past_key_values") is not None:
119
+ key_states, value_states = kwargs["past_key_values"].update(
120
+ key_states, value_states, self.layer_idx
121
+ )
122
+
123
+ if ALL_ATTENTION_FUNCTIONS is None:
124
+ attention_interface = eager_attention_forward
125
+ elif hasattr(ALL_ATTENTION_FUNCTIONS, "get_interface"):
126
+ attention_interface = ALL_ATTENTION_FUNCTIONS.get_interface(
127
+ self.config._attn_implementation, eager_attention_forward
128
+ )
129
+ else: # pragma: no cover - compatibility with older Transformers.
130
+ attention_interface = ALL_ATTENTION_FUNCTIONS.get(
131
+ self.config._attn_implementation, eager_attention_forward
132
+ )
133
+ is_causal = attention_mask is None and key_states.shape[-2] == q.shape[1]
134
+ y, _ = attention_interface(
135
+ self,
136
+ q.transpose(1, 2),
137
+ key_states,
138
+ value_states,
139
+ attention_mask,
140
+ is_causal=is_causal,
141
+ **kwargs,
142
+ )
143
+ y = y.contiguous().view_as(x)
144
+ return self.attn_resid(y)
145
+
146
+
147
+ class MLP(nn.Module):
148
+ def __init__(self, config: TalkieConfig):
149
+ super().__init__()
150
+ n_state = config.n_embd
151
+ n_mlp = int(round(((8 / 3) * n_state) / 128) * 128)
152
+
153
+ self.mlp_gate = nn.Linear(n_state, n_mlp, bias=False)
154
+ self.mlp_linear = nn.Linear(n_state, n_mlp, bias=False)
155
+ self.mlp_resid = nn.Linear(n_mlp, n_state, bias=False)
156
+
157
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
158
+ x = F.silu(self.mlp_gate(x)) * self.mlp_linear(x)
159
+ return self.mlp_resid(x)
160
+
161
+
162
+ class Block(nn.Module):
163
+ def __init__(self, config: TalkieConfig, layer_idx: int):
164
+ super().__init__()
165
+ self.attn = CausalSelfAttention(config, layer_idx)
166
+ self.attn_gain = ActGain((2 * config.n_layer) ** -0.5)
167
+ self.mlp = MLP(config)
168
+ self.mlp_gain = ActGain((2 * config.n_layer) ** -0.5)
169
+ self.embed_skip = ActGain(0.0)
170
+
171
+ def forward(
172
+ self,
173
+ e_x: torch.Tensor,
174
+ x: torch.Tensor,
175
+ cos_sin: tuple[torch.Tensor, torch.Tensor],
176
+ attention_mask: torch.Tensor | None = None,
177
+ **kwargs,
178
+ ) -> torch.Tensor:
179
+ x = x + self.attn_gain(
180
+ self.attn(F.rms_norm(x, (x.shape[-1],)), cos_sin, attention_mask, **kwargs)
181
+ )
182
+ x = x + self.mlp_gain(self.mlp(F.rms_norm(x, (x.shape[-1],))))
183
+ x = x + self.embed_skip(e_x)
184
+ return x
185
+
186
+
187
+ class TalkiePreTrainedModel(PreTrainedModel):
188
+ config_class = TalkieConfig
189
+ base_model_prefix = ""
190
+ supports_gradient_checkpointing = True
191
+ _supports_sdpa = True
192
+ _supports_attention_backend = True
193
+ _no_split_modules = ["Block"]
194
+ _tied_weights_keys = None
195
+
196
+ def _init_weights(self, module: nn.Module) -> None:
197
+ return
198
+
199
+
200
+ class TalkieModel(TalkiePreTrainedModel, GenerationMixin):
201
+ def __init__(self, config: TalkieConfig):
202
+ super().__init__(config)
203
+ self.embed = nn.Embedding(config.vocab_size, config.n_embd)
204
+ self.blocks = nn.ModuleList([Block(config, i) for i in range(config.n_layer)])
205
+ self.gradient_checkpointing = False
206
+
207
+ cos, sin = self._precompute_rotary_embeddings(config.max_position_embeddings)
208
+ self.register_buffer("cos", cos, persistent=False)
209
+ self.register_buffer("sin", sin, persistent=False)
210
+ self._rotary_initialized = cos.device.type != "meta"
211
+ self.post_init()
212
+
213
+ def _precompute_rotary_embeddings(
214
+ self,
215
+ seq_len: int,
216
+ head_dim: int | None = None,
217
+ base: int | float | None = None,
218
+ ) -> tuple[torch.Tensor, torch.Tensor]:
219
+ device = self.embed.weight.device if hasattr(self, "embed") else "cpu"
220
+ head_dim = head_dim if head_dim is not None else self.config.head_dim
221
+ base = base if base is not None else self.config.rope_base
222
+ inv_freq, attention_factor = self._rotary_inv_freq(seq_len, head_dim, float(base), device)
223
+ t = torch.arange(seq_len, dtype=torch.float32, device=device)
224
+ freqs = torch.outer(t, inv_freq)
225
+ cos, sin = freqs.cos(), freqs.sin()
226
+ if attention_factor != 1.0:
227
+ cos = cos * attention_factor
228
+ sin = sin * attention_factor
229
+ cos, sin = cos.bfloat16(), sin.bfloat16()
230
+ cos, sin = cos[None, :, None, :], sin[None, :, None, :]
231
+ return cos, sin
232
+
233
+ def _rotary_inv_freq(
234
+ self,
235
+ seq_len: int,
236
+ head_dim: int,
237
+ base: float,
238
+ device: torch.device | str,
239
+ ) -> tuple[torch.Tensor, float]:
240
+ scaling = self.config.rope_scaling
241
+ rope_type = scaling.get("rope_type") if scaling else None
242
+ if rope_type in (None, "default"):
243
+ return self._default_rotary_inv_freq(head_dim, base, device), 1.0
244
+ if rope_type == "linear":
245
+ inv_freq = self._default_rotary_inv_freq(head_dim, base, device)
246
+ return inv_freq / float(scaling["factor"]), 1.0
247
+ if rope_type == "dynamic":
248
+ return self._dynamic_rotary_inv_freq(seq_len, head_dim, base, device, scaling), 1.0
249
+ if rope_type == "yarn":
250
+ return self._yarn_rotary_inv_freq(head_dim, base, device, scaling)
251
+ raise ValueError(f"unsupported rope_scaling type {rope_type!r}")
252
+
253
+ @staticmethod
254
+ def _default_rotary_inv_freq(
255
+ head_dim: int, base: float, device: torch.device | str
256
+ ) -> torch.Tensor:
257
+ channel_range = torch.arange(0, head_dim, 2, dtype=torch.float32, device=device)
258
+ return 1.0 / (base ** (channel_range / head_dim))
259
+
260
+ def _original_max_position_embeddings(self, scaling: dict | None) -> int:
261
+ if scaling and "original_max_position_embeddings" in scaling:
262
+ return int(scaling["original_max_position_embeddings"])
263
+ return int(self.config.max_position_embeddings)
264
+
265
+ def _dynamic_rotary_inv_freq(
266
+ self,
267
+ seq_len: int,
268
+ head_dim: int,
269
+ base: float,
270
+ device: torch.device | str,
271
+ scaling: dict,
272
+ ) -> torch.Tensor:
273
+ original_max_position_embeddings = self._original_max_position_embeddings(scaling)
274
+ scaled_seq_len = max(seq_len, original_max_position_embeddings)
275
+ factor = float(scaling["factor"])
276
+ base = base * (
277
+ (factor * scaled_seq_len / original_max_position_embeddings) - (factor - 1.0)
278
+ ) ** (head_dim / (head_dim - 2.0))
279
+ return self._default_rotary_inv_freq(head_dim, base, device)
280
+
281
+ def _yarn_rotary_inv_freq(
282
+ self,
283
+ head_dim: int,
284
+ base: float,
285
+ device: torch.device | str,
286
+ scaling: dict,
287
+ ) -> tuple[torch.Tensor, float]:
288
+ factor = float(scaling["factor"])
289
+ original_max_position_embeddings = self._original_max_position_embeddings(scaling)
290
+ beta_fast = float(scaling.get("beta_fast", 32.0))
291
+ beta_slow = float(scaling.get("beta_slow", 1.0))
292
+ attention_factor = scaling.get("attention_factor")
293
+ if attention_factor is None:
294
+ attention_factor = 1.0 if factor <= 1.0 else 0.1 * math.log(factor) + 1.0
295
+
296
+ channel_range = torch.arange(0, head_dim, 2, dtype=torch.float32, device=device)
297
+ pos_freqs = base ** (channel_range / head_dim)
298
+ inv_freq_extrapolation = 1.0 / pos_freqs
299
+ inv_freq_interpolation = 1.0 / (factor * pos_freqs)
300
+
301
+ low, high = self._yarn_correction_range(
302
+ beta_fast,
303
+ beta_slow,
304
+ head_dim,
305
+ base,
306
+ original_max_position_embeddings,
307
+ truncate=bool(scaling.get("truncate", True)),
308
+ )
309
+ ramp = self._yarn_linear_ramp(low, high, head_dim // 2, device)
310
+ extrapolation_factor = 1.0 - ramp
311
+ inv_freq = (
312
+ inv_freq_interpolation * (1.0 - extrapolation_factor)
313
+ + inv_freq_extrapolation * extrapolation_factor
314
+ )
315
+ return inv_freq, float(attention_factor)
316
+
317
+ @staticmethod
318
+ def _yarn_correction_range(
319
+ low_rot: float,
320
+ high_rot: float,
321
+ head_dim: int,
322
+ base: float,
323
+ original_max_position_embeddings: int,
324
+ truncate: bool,
325
+ ) -> tuple[float, float]:
326
+ def correction_dim(num_rotations: float) -> float:
327
+ return (
328
+ head_dim
329
+ * math.log(original_max_position_embeddings / (num_rotations * 2.0 * math.pi))
330
+ / (2.0 * math.log(base))
331
+ )
332
+
333
+ low = correction_dim(low_rot)
334
+ high = correction_dim(high_rot)
335
+ if truncate:
336
+ low = math.floor(low)
337
+ high = math.ceil(high)
338
+ return max(low, 0.0), min(high, float(head_dim - 1))
339
+
340
+ @staticmethod
341
+ def _yarn_linear_ramp(
342
+ low: float,
343
+ high: float,
344
+ dim: int,
345
+ device: torch.device | str,
346
+ ) -> torch.Tensor:
347
+ if low == high:
348
+ high += 0.001
349
+ ramp = (torch.arange(dim, dtype=torch.float32, device=device) - low) / (high - low)
350
+ return torch.clamp(ramp, 0.0, 1.0)
351
+
352
+ def _ensure_rotary_embeddings(self, seq_len: int) -> None:
353
+ device = self.embed.weight.device
354
+ needs_init = (
355
+ not self._rotary_initialized
356
+ or self.cos.device != device
357
+ or self.sin.device != device
358
+ or self.cos.shape[1] < seq_len
359
+ )
360
+ if needs_init:
361
+ max_seq_len = max(seq_len, self.config.max_position_embeddings)
362
+ cos, sin = self._precompute_rotary_embeddings(max_seq_len)
363
+ self.cos = cos.to(device=device)
364
+ self.sin = sin.to(device=device)
365
+ self._rotary_initialized = True
366
+
367
+ def reset_rotary_embeddings(self) -> None:
368
+ self._rotary_initialized = False
369
+
370
+ def get_input_embeddings(self) -> nn.Embedding:
371
+ return self.embed
372
+
373
+ def set_input_embeddings(self, value: nn.Embedding) -> None:
374
+ self.embed = value
375
+
376
+ def _position_ids(
377
+ self,
378
+ input_ids: torch.LongTensor,
379
+ position_ids: torch.LongTensor | None = None,
380
+ cache_position: torch.LongTensor | None = None,
381
+ past_key_values: Cache | None = None,
382
+ ) -> torch.LongTensor:
383
+ batch_size, seq_len = input_ids.shape
384
+ if position_ids is not None:
385
+ if position_ids.dim() == 1:
386
+ position_ids = position_ids.unsqueeze(0)
387
+ return position_ids.to(device=input_ids.device, dtype=torch.long)
388
+ if cache_position is not None:
389
+ if cache_position.dim() == 1:
390
+ cache_position = cache_position.unsqueeze(0)
391
+ if cache_position.shape[0] == 1 and batch_size != 1:
392
+ cache_position = cache_position.expand(batch_size, -1)
393
+ return cache_position.to(device=input_ids.device, dtype=torch.long)
394
+ past_seen = past_key_values.get_seq_length() if past_key_values is not None else 0
395
+ position_ids = torch.arange(seq_len, device=input_ids.device, dtype=torch.long) + past_seen
396
+ return position_ids.unsqueeze(0).expand(batch_size, -1)
397
+
398
+ def _attention_mask(
399
+ self,
400
+ attention_mask: torch.Tensor | None,
401
+ input_ids: torch.Tensor,
402
+ position_ids: torch.Tensor,
403
+ past_key_values: Cache | None,
404
+ dtype: torch.dtype,
405
+ ) -> torch.Tensor | None:
406
+ if attention_mask is not None and attention_mask.dim() >= 4:
407
+ return attention_mask
408
+ batch_size, query_length = input_ids.shape
409
+ past_seen = past_key_values.get_seq_length() if past_key_values is not None else 0
410
+
411
+ if attention_mask is not None and attention_mask.dim() != 2:
412
+ return attention_mask
413
+ if attention_mask is None and past_seen == 0:
414
+ return None
415
+
416
+ key_length = past_seen + query_length
417
+ if attention_mask is not None:
418
+ if attention_mask.shape[-1] == query_length and past_seen:
419
+ prefix = torch.ones(
420
+ attention_mask.shape[0],
421
+ past_seen,
422
+ dtype=attention_mask.dtype,
423
+ device=attention_mask.device,
424
+ )
425
+ attention_mask = torch.cat([prefix, attention_mask], dim=-1)
426
+ key_length = attention_mask.shape[-1]
427
+
428
+ key_positions = torch.arange(key_length, device=input_ids.device, dtype=torch.long)
429
+ future_mask = key_positions.view(1, 1, 1, key_length) > position_ids.view(
430
+ batch_size, 1, query_length, 1
431
+ )
432
+ if attention_mask is not None:
433
+ padding_mask = attention_mask[:, None, None, :].to(device=input_ids.device) == 0
434
+ mask = future_mask | padding_mask
435
+ else:
436
+ mask = future_mask
437
+
438
+ min_value = torch.finfo(dtype).min
439
+ causal_mask = torch.zeros(
440
+ batch_size, 1, query_length, key_length, dtype=dtype, device=input_ids.device
441
+ )
442
+ return causal_mask.masked_fill(mask, min_value)
443
+
444
+ def forward(
445
+ self,
446
+ input_ids: torch.LongTensor | None = None,
447
+ inputs_embeds: torch.FloatTensor | None = None,
448
+ attention_mask: torch.Tensor | None = None,
449
+ position_ids: torch.LongTensor | None = None,
450
+ past_key_values: Cache | None = None,
451
+ use_cache: bool | None = None,
452
+ return_dict: bool | None = None,
453
+ **kwargs,
454
+ ) -> BaseModelOutputWithPast | tuple[torch.Tensor, ...]:
455
+ cache_position = kwargs.pop("cache_position", None)
456
+ if input_ids is None and inputs_embeds is None:
457
+ raise ValueError("input_ids or inputs_embeds is required")
458
+ if input_ids is not None and inputs_embeds is not None:
459
+ raise ValueError("provide only one of input_ids or inputs_embeds")
460
+ if input_ids is None:
461
+ input_ids = torch.empty(
462
+ inputs_embeds.shape[:2],
463
+ dtype=torch.long,
464
+ device=inputs_embeds.device,
465
+ )
466
+ use_cache = use_cache if use_cache is not None else self.config.use_cache
467
+ if self.gradient_checkpointing and self.training:
468
+ use_cache = False
469
+ if use_cache and past_key_values is None:
470
+ past_key_values = DynamicCache(config=self.config)
471
+
472
+ position_ids = self._position_ids(input_ids, position_ids, cache_position, past_key_values)
473
+ # Keep graph capture free of CUDA tensor -> Python scalar syncs. The
474
+ # configured context length is the static serving/training contract.
475
+ self._ensure_rotary_embeddings(int(self.config.max_position_embeddings))
476
+
477
+ cos = self.cos[0, position_ids, :, :]
478
+ sin = self.sin[0, position_ids, :, :]
479
+ cos_sin = cos, sin
480
+ x = inputs_embeds if inputs_embeds is not None else self.embed(input_ids)
481
+ x = F.rms_norm(x, (x.shape[-1],))
482
+ attention_mask = self._attention_mask(attention_mask, input_ids, position_ids, past_key_values, x.dtype)
483
+ e_x = x
484
+ for block in self.blocks:
485
+ if self.gradient_checkpointing and self.training:
486
+ def custom_forward(
487
+ e_x: torch.Tensor,
488
+ x: torch.Tensor,
489
+ cos: torch.Tensor,
490
+ sin: torch.Tensor,
491
+ attention_mask: torch.Tensor | None,
492
+ block: Block = block,
493
+ ) -> torch.Tensor:
494
+ return block(e_x, x, (cos, sin), attention_mask=attention_mask)
495
+
496
+ x = self._gradient_checkpointing_func(
497
+ custom_forward,
498
+ e_x,
499
+ x,
500
+ cos,
501
+ sin,
502
+ attention_mask,
503
+ )
504
+ else:
505
+ x = block(
506
+ e_x,
507
+ x,
508
+ cos_sin,
509
+ attention_mask=attention_mask,
510
+ past_key_values=past_key_values if use_cache else None,
511
+ **kwargs,
512
+ )
513
+ x = F.rms_norm(x, (x.shape[-1],))
514
+ past_key_values = past_key_values if use_cache else None
515
+ use_return_dict = return_dict if return_dict is not None else self.config.use_return_dict
516
+ if use_return_dict:
517
+ return BaseModelOutputWithPast(last_hidden_state=x, past_key_values=past_key_values)
518
+ output = (x,)
519
+ return output + ((past_key_values,) if past_key_values is not None else ())
520
+
521
+
522
+ class TalkieForCausalLM(TalkieModel):
523
+ _tied_weights_keys = None
524
+
525
+ def __init__(self, config: TalkieConfig):
526
+ super().__init__(config)
527
+ self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
528
+ self.post_init()
529
+
530
+ def get_output_embeddings(self) -> nn.Linear:
531
+ return self.lm_head
532
+
533
+ def set_output_embeddings(self, value: nn.Linear) -> None:
534
+ self.lm_head = value
535
+
536
+ def _chunked_lm_loss(
537
+ self,
538
+ hidden_states: torch.Tensor,
539
+ labels: torch.Tensor,
540
+ chunk_size: int,
541
+ ) -> torch.Tensor:
542
+ if chunk_size <= 0:
543
+ raise ValueError("chunk_size must be positive")
544
+
545
+ total_loss = hidden_states.new_zeros((), dtype=torch.float32)
546
+ total_tokens = hidden_states.new_zeros((), dtype=torch.float32)
547
+ for start in range(0, hidden_states.shape[1], chunk_size):
548
+ end = min(start + chunk_size, hidden_states.shape[1])
549
+ logits = self.lm_head(hidden_states[:, start:end, :]).float()
550
+ if self.config.logit_scale != 1.0:
551
+ logits = logits * self.config.logit_scale
552
+ chunk_labels = labels[:, start:end].contiguous()
553
+ total_loss = total_loss + F.cross_entropy(
554
+ logits.reshape(-1, logits.size(-1)),
555
+ chunk_labels.reshape(-1),
556
+ ignore_index=-100,
557
+ reduction="sum",
558
+ )
559
+ total_tokens = total_tokens + (chunk_labels != -100).sum(dtype=torch.float32)
560
+ return total_loss / total_tokens.clamp_min(1.0)
561
+
562
+ def forward(
563
+ self,
564
+ input_ids: torch.LongTensor | None = None,
565
+ attention_mask: torch.Tensor | None = None,
566
+ inputs_embeds: torch.FloatTensor | None = None,
567
+ labels: torch.LongTensor | None = None,
568
+ return_dict: bool | None = None,
569
+ past_key_values: Cache | None = None,
570
+ use_cache: bool | None = None,
571
+ position_ids: torch.LongTensor | None = None,
572
+ logits_to_keep: int | torch.Tensor = 0,
573
+ loss_chunk_size: int = 0,
574
+ return_logits: bool = True,
575
+ **kwargs,
576
+ ) -> CausalLMOutputWithPast | tuple[torch.Tensor, ...]:
577
+ if input_ids is None and inputs_embeds is None:
578
+ raise ValueError("input_ids or inputs_embeds is required")
579
+ cache_position = kwargs.pop("cache_position", None)
580
+ outputs = super().forward(
581
+ input_ids,
582
+ inputs_embeds=inputs_embeds,
583
+ attention_mask=attention_mask,
584
+ position_ids=position_ids,
585
+ past_key_values=past_key_values,
586
+ use_cache=use_cache,
587
+ cache_position=cache_position,
588
+ return_dict=True,
589
+ **kwargs,
590
+ )
591
+ hidden_states = outputs.last_hidden_state
592
+ loss = None
593
+ logits = None
594
+ if labels is not None and loss_chunk_size > 0:
595
+ loss = self._chunked_lm_loss(
596
+ hidden_states[:, :-1, :],
597
+ labels[:, 1:],
598
+ loss_chunk_size,
599
+ )
600
+ if return_logits:
601
+ slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep
602
+ logits = self.lm_head(hidden_states[:, slice_indices, :]).float()
603
+ if self.config.logit_scale != 1.0:
604
+ logits = logits * self.config.logit_scale
605
+
606
+ if labels is not None and loss is None:
607
+ if logits is None:
608
+ raise ValueError("return_logits must be true when loss_chunk_size is not used")
609
+ shift_logits = logits[..., :-1, :].contiguous()
610
+ shift_labels = labels[..., 1:].contiguous()
611
+ loss = F.cross_entropy(
612
+ shift_logits.view(-1, shift_logits.size(-1)),
613
+ shift_labels.view(-1),
614
+ ignore_index=-100,
615
+ )
616
+
617
+ use_return_dict = return_dict if return_dict is not None else self.config.use_return_dict
618
+ if use_return_dict:
619
+ return CausalLMOutputWithPast(
620
+ loss=loss,
621
+ logits=logits,
622
+ past_key_values=outputs.past_key_values,
623
+ )
624
+ output = (logits,)
625
+ if outputs.past_key_values is not None:
626
+ output += (outputs.past_key_values,)
627
+ return ((loss,) + output) if loss is not None else output
628
+
629
+
630
+ __all__ = ["TalkieConfig", "TalkieForCausalLM", "TalkieModel"]
tokenization_talkie.py ADDED
@@ -0,0 +1,168 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import os
4
+ import shutil
5
+ from pathlib import Path
6
+
7
+ import tiktoken
8
+ from tiktoken.load import load_tiktoken_bpe
9
+ from transformers import PreTrainedTokenizer
10
+
11
+
12
+ BASE_VOCAB_SIZE = 65536
13
+ IT_VOCAB_SIZE = BASE_VOCAB_SIZE + 4
14
+
15
+ _PAT_STR = "|".join(
16
+ [
17
+ r"""[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}]*[\p{Ll}\p{Lm}\p{Lo}\p{M}]+(?i:'s|'t|'re|'ve|'m|'ll|'d)?""",
18
+ r"""[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}]+[\p{Ll}\p{Lm}\p{Lo}\p{M}]*(?i:'s|'t|'re|'ve|'m|'ll|'d)?""",
19
+ r"""\p{N}{1,3}""",
20
+ r""" ?[^\s\p{L}\p{N}]+[\r\n/]*""",
21
+ r"""\s*[\r\n]+""",
22
+ r"""\s+(?!\S)""",
23
+ r"""\s+""",
24
+ ]
25
+ )
26
+
27
+ _BASE_SPECIAL_TOKENS = {
28
+ "<|endoftext|>": BASE_VOCAB_SIZE - 1,
29
+ }
30
+
31
+ _IT_SPECIAL_TOKENS = {
32
+ "<|endoftext|>": BASE_VOCAB_SIZE - 1,
33
+ "<|end|>": BASE_VOCAB_SIZE,
34
+ "<|user|>": BASE_VOCAB_SIZE + 1,
35
+ "<|assistant|>": BASE_VOCAB_SIZE + 2,
36
+ "<|system|>": BASE_VOCAB_SIZE + 3,
37
+ }
38
+
39
+
40
+ class TalkieTokenizer(PreTrainedTokenizer):
41
+ vocab_files_names = {"vocab_file": "vocab.txt"}
42
+ model_input_names = ["input_ids", "attention_mask"]
43
+
44
+ def __init__(
45
+ self,
46
+ vocab_file: str,
47
+ style: str = "base",
48
+ model_max_length: int = 2048,
49
+ **kwargs,
50
+ ):
51
+ self.vocab_file = str(vocab_file)
52
+ self.style = style
53
+
54
+ mergeable_ranks = load_tiktoken_bpe(self.vocab_file)
55
+ mergeable_ranks = {
56
+ key: value for key, value in mergeable_ranks.items() if value < BASE_VOCAB_SIZE - 1
57
+ }
58
+ if style == "it":
59
+ special_tokens = dict(_IT_SPECIAL_TOKENS)
60
+ vocab_size = IT_VOCAB_SIZE
61
+ name = "talkie-it"
62
+ elif style == "base":
63
+ special_tokens = dict(_BASE_SPECIAL_TOKENS)
64
+ vocab_size = BASE_VOCAB_SIZE
65
+ name = "talkie-base"
66
+ else:
67
+ raise ValueError(f"unknown Talkie tokenizer style: {style!r}")
68
+
69
+ self.encoder = tiktoken.Encoding(
70
+ name=name,
71
+ pat_str=_PAT_STR,
72
+ mergeable_ranks=mergeable_ranks,
73
+ special_tokens=special_tokens,
74
+ )
75
+ self._vocab_size = vocab_size
76
+ self._special_token_to_id = special_tokens
77
+ self._id_to_special_token = {value: key for key, value in special_tokens.items()}
78
+
79
+ if style == "it":
80
+ kwargs.setdefault("eos_token", "<|end|>")
81
+ kwargs.setdefault(
82
+ "additional_special_tokens",
83
+ ["<|endoftext|>", "<|user|>", "<|assistant|>", "<|system|>"],
84
+ )
85
+ else:
86
+ kwargs.setdefault("eos_token", "<|endoftext|>")
87
+ super().__init__(model_max_length=model_max_length, **kwargs)
88
+
89
+ @property
90
+ def vocab_size(self) -> int:
91
+ return self._vocab_size
92
+
93
+ def get_vocab(self) -> dict[str, int]:
94
+ vocab = {str(index): index for index in range(self._vocab_size)}
95
+ vocab.update(self._special_token_to_id)
96
+ vocab.update(self.get_added_vocab())
97
+ return vocab
98
+
99
+ def _tokenize(self, text: str, **kwargs) -> list[str]:
100
+ return [str(token_id) for token_id in self.encoder.encode(text, allowed_special="all")]
101
+
102
+ def _convert_token_to_id(self, token: str) -> int:
103
+ if token in self._special_token_to_id:
104
+ return self._special_token_to_id[token]
105
+ try:
106
+ token_id = int(token)
107
+ except ValueError:
108
+ return self.eos_token_id
109
+ if 0 <= token_id < self._vocab_size:
110
+ return token_id
111
+ return self.eos_token_id
112
+
113
+ def _convert_id_to_token(self, index: int) -> str:
114
+ index = int(index)
115
+ return self._id_to_special_token.get(index, str(index))
116
+
117
+ def convert_tokens_to_string(self, tokens: list[str]) -> str:
118
+ ids = [self._convert_token_to_id(token) for token in tokens]
119
+ return self.encoder.decode(ids)
120
+
121
+ def _decode(
122
+ self,
123
+ token_ids,
124
+ skip_special_tokens: bool = False,
125
+ clean_up_tokenization_spaces: bool | None = None,
126
+ **kwargs,
127
+ ) -> str:
128
+ if isinstance(token_ids, int):
129
+ token_ids = [token_ids]
130
+ ids = [int(token_id) for token_id in token_ids]
131
+ if skip_special_tokens:
132
+ specials = set(self._special_token_to_id.values())
133
+ ids = [token_id for token_id in ids if token_id not in specials]
134
+ return self.encoder.decode(ids)
135
+
136
+ def build_inputs_with_special_tokens(
137
+ self, token_ids_0: list[int], token_ids_1: list[int] | None = None
138
+ ) -> list[int]:
139
+ if token_ids_1 is None:
140
+ return list(token_ids_0)
141
+ return list(token_ids_0) + list(token_ids_1)
142
+
143
+ def get_special_tokens_mask(
144
+ self,
145
+ token_ids_0: list[int],
146
+ token_ids_1: list[int] | None = None,
147
+ already_has_special_tokens: bool = False,
148
+ ) -> list[int]:
149
+ special_ids = set(self._special_token_to_id.values())
150
+ if already_has_special_tokens:
151
+ return [1 if token_id in special_ids else 0 for token_id in token_ids_0]
152
+ token_ids = list(token_ids_0) if token_ids_1 is None else list(token_ids_0) + list(token_ids_1)
153
+ return [1 if token_id in special_ids else 0 for token_id in token_ids]
154
+
155
+ def create_token_type_ids_from_sequences(
156
+ self, token_ids_0: list[int], token_ids_1: list[int] | None = None
157
+ ) -> list[int]:
158
+ length = len(token_ids_0) if token_ids_1 is None else len(token_ids_0) + len(token_ids_1)
159
+ return [0] * length
160
+
161
+ def save_vocabulary(self, save_directory: str, filename_prefix: str | None = None):
162
+ if not os.path.isdir(save_directory):
163
+ raise ValueError(f"Vocabulary path {save_directory!r} is not a directory")
164
+ name = "vocab.txt" if filename_prefix is None else f"{filename_prefix}-vocab.txt"
165
+ out = Path(save_directory) / name
166
+ if Path(self.vocab_file).resolve() != out.resolve():
167
+ shutil.copyfile(self.vocab_file, out)
168
+ return (str(out),)
tokenizer_config.json ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "added_tokens_decoder": {
3
+ "65535": {
4
+ "content": "<|endoftext|>",
5
+ "lstrip": false,
6
+ "normalized": false,
7
+ "rstrip": false,
8
+ "single_word": false,
9
+ "special": true
10
+ }
11
+ },
12
+ "auto_map": {
13
+ "AutoTokenizer": [
14
+ "tokenization_talkie.TalkieTokenizer",
15
+ null
16
+ ]
17
+ },
18
+ "backend": "custom",
19
+ "eos_token": "<|endoftext|>",
20
+ "is_local": true,
21
+ "local_files_only": false,
22
+ "model_max_length": 9223372036854775807,
23
+ "tokenizer_class": "TalkieTokenizer"
24
+ }
vocab.txt ADDED
The diff for this file is too large to render. See raw diff