atrost commited on
Commit
8cc112e
·
verified ·
1 Parent(s): a24f09a

Upload vLLM-compatible full-width draft conversion from 1ed51485

Browse files
README.md ADDED
@@ -0,0 +1,50 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: cc-by-nc-4.0
3
+ language:
4
+ - en
5
+ datasets:
6
+ - nvidia/Nemotron-ClimbMix
7
+ tags:
8
+ - stairformer
9
+ - asymmetric
10
+ - vllm
11
+ - causal-lm
12
+ - pretraining
13
+ - climbmix
14
+ - custom_code
15
+ ---
16
+
17
+ # atrost/climbmix-stairformer-353m-extracted-nested-94m-1p2b-h100-vllm
18
+
19
+ vLLM compatibility conversion of `atrost/climbmix-stairformer-353m-extracted-nested-94m-1p2b-h100`.
20
+
21
+ The source checkpoint is a compact asymmetric draft model: its entry layer uses
22
+ the large 1280-wide embedding space, then later layers and the LM head operate in
23
+ the 160-wide nested space. vLLM's generic Transformers causal wrapper assumes one
24
+ hidden size for the base model and LM head, so this repo expands the checkpoint
25
+ into a full-width StairFormer-shaped model with masked/zeroed suffix weights.
26
+
27
+ The converted model preserves the source logits up to floating-point differences,
28
+ but it is a compatibility artifact rather than an optimized 94M draft runtime.
29
+
30
+ - Source checkpoint: `atrost/climbmix-stairformer-353m-extracted-nested-94m-1p2b-h100`
31
+ - Source revision: `1ed51485d6ed35267425e46c77bd9daf86b781aa`
32
+ - Local parity max logits diff: `4.292e-06`
33
+
34
+ ```python
35
+ from transformers import AutoModelForCausalLM, AutoTokenizer
36
+
37
+ repo_id = "atrost/climbmix-stairformer-353m-extracted-nested-94m-1p2b-h100-vllm"
38
+ tokenizer = AutoTokenizer.from_pretrained(repo_id)
39
+ model = AutoModelForCausalLM.from_pretrained(
40
+ repo_id,
41
+ trust_remote_code=True,
42
+ torch_dtype="auto",
43
+ )
44
+ ```
45
+
46
+ ```python
47
+ from vllm import LLM
48
+
49
+ llm = LLM(model="atrost/climbmix-stairformer-353m-extracted-nested-94m-1p2b-h100-vllm", trust_remote_code=True, model_impl="transformers")
50
+ ```
config.json ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "StairFormerForCausalLM"
4
+ ],
5
+ "attention_bias": false,
6
+ "attention_dropout": 0.0,
7
+ "auto_map": {
8
+ "AutoModel": "modeling_stairformer.StairFormerModel",
9
+ "AutoModelForCausalLM": "modeling_stairformer.StairFormerForCausalLM"
10
+ },
11
+ "bos_token_id": 50256,
12
+ "dtype": "float32",
13
+ "eos_token_id": 50256,
14
+ "head_dim": 80,
15
+ "hidden_act": "silu",
16
+ "hidden_size": 1280,
17
+ "initializer_range": 0.02,
18
+ "intermediate_size": 3584,
19
+ "max_position_embeddings": 2048,
20
+ "mlp_bias": false,
21
+ "model_type": "llama",
22
+ "num_attention_heads": 16,
23
+ "num_hidden_layers": 12,
24
+ "num_key_value_heads": 8,
25
+ "pad_token_id": 50256,
26
+ "pretraining_tp": 1,
27
+ "rms_norm_eps": 1e-05,
28
+ "rope_parameters": {
29
+ "rope_theta": 10000.0,
30
+ "rope_type": "default"
31
+ },
32
+ "stairformer_dense_entry_layers": 1,
33
+ "stairformer_nested_loss_alpha": 0.1,
34
+ "stairformer_prefix_hidden_size": 160,
35
+ "stairformer_prefix_intermediate_size": 448,
36
+ "tie_word_embeddings": false,
37
+ "transformers_version": "5.7.0",
38
+ "use_cache": true,
39
+ "vocab_size": 50304
40
+ }
generation_config.json ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "bos_token_id": 50256,
4
+ "eos_token_id": 50256,
5
+ "output_attentions": false,
6
+ "output_hidden_states": false,
7
+ "pad_token_id": 50256,
8
+ "transformers_version": "5.7.0",
9
+ "use_cache": true
10
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:08299b6749f589bb52531f532cdd4c57336d3857a00b96511de114e5114d6040
3
+ size 1617250624
modeling_stairformer.py ADDED
@@ -0,0 +1,579 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Remote-code definitions for StairFormer and asymmetric nested Llama checkpoints."""
3
+
4
+ from __future__ import annotations
5
+
6
+ from dataclasses import dataclass
7
+ from typing import List, Optional, Tuple, Union
8
+
9
+ import torch
10
+ import torch.nn as nn
11
+ import torch.nn.functional as F
12
+ from torch.nn import CrossEntropyLoss
13
+
14
+ from transformers import LlamaConfig, LlamaForCausalLM, LlamaPreTrainedModel
15
+ from transformers.generation import GenerationMixin
16
+ from transformers.modeling_outputs import BaseModelOutputWithPast, ModelOutput
17
+ from transformers.models.llama.modeling_llama import (
18
+ LlamaDecoderLayer,
19
+ LlamaModel,
20
+ LlamaRMSNorm,
21
+ LlamaRotaryEmbedding,
22
+ )
23
+
24
+
25
+ @dataclass
26
+ class StairFormerCausalLMOutputWithPast(ModelOutput):
27
+ loss: Optional[torch.FloatTensor] = None
28
+ loss_full: Optional[torch.FloatTensor] = None
29
+ loss_nested: Optional[torch.FloatTensor] = None
30
+ logits: Optional[torch.FloatTensor] = None
31
+ nested_logits: Optional[torch.FloatTensor] = None
32
+ past_key_values: Optional[Tuple[Tuple[torch.FloatTensor]]] = None
33
+ hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
34
+ attentions: Optional[Tuple[torch.FloatTensor, ...]] = None
35
+
36
+
37
+ @dataclass
38
+ class AsymmetricNestedLlamaOutput(ModelOutput):
39
+ loss: Optional[torch.FloatTensor] = None
40
+ logits: torch.FloatTensor = None
41
+ hidden_states: torch.FloatTensor = None
42
+ past_key_values: Optional[Tuple[Tuple[torch.FloatTensor]]] = None
43
+
44
+
45
+ def configure_stairformer_config(config: LlamaConfig) -> LlamaConfig:
46
+ if getattr(config, "_attn_implementation", None) is None:
47
+ config._attn_implementation = "eager"
48
+ head_dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)
49
+ query_heads_per_kv = config.num_attention_heads // config.num_key_value_heads
50
+ config.stairformer_prefix_hidden_size = getattr(
51
+ config,
52
+ "stairformer_prefix_hidden_size",
53
+ head_dim * query_heads_per_kv,
54
+ )
55
+ config.stairformer_prefix_intermediate_size = getattr(
56
+ config,
57
+ "stairformer_prefix_intermediate_size",
58
+ config.intermediate_size * config.stairformer_prefix_hidden_size // config.hidden_size,
59
+ )
60
+ config.stairformer_dense_entry_layers = getattr(config, "stairformer_dense_entry_layers", 1)
61
+ config.stairformer_nested_loss_alpha = getattr(
62
+ config,
63
+ "stairformer_nested_loss_alpha",
64
+ getattr(config, "stairformer_nested_loss_weight", 0.1),
65
+ )
66
+ return config
67
+
68
+
69
+ class BlockLowerTriangularLinear(nn.Module):
70
+ def __init__(
71
+ self,
72
+ in_features: int,
73
+ out_features: int,
74
+ prefix_in_features: int,
75
+ prefix_out_features: int,
76
+ bias: bool = False,
77
+ device=None,
78
+ dtype=None,
79
+ ):
80
+ super().__init__()
81
+ factory_kwargs = {"device": device, "dtype": dtype}
82
+ self.in_features = in_features
83
+ self.out_features = out_features
84
+ self.prefix_in_features = prefix_in_features
85
+ self.prefix_out_features = prefix_out_features
86
+ self.weight = nn.Parameter(torch.empty(out_features, in_features, **factory_kwargs))
87
+ self.bias = nn.Parameter(torch.empty(out_features, **factory_kwargs)) if bias else None
88
+ mask = torch.ones(out_features, in_features, device=device, dtype=torch.bool)
89
+ mask[:prefix_out_features, prefix_in_features:] = False
90
+ self.register_buffer("weight_mask", mask)
91
+ self.reset_parameters()
92
+
93
+ @classmethod
94
+ def from_linear(
95
+ cls,
96
+ linear: nn.Linear,
97
+ prefix_in_features: int,
98
+ prefix_out_features: int,
99
+ ) -> "BlockLowerTriangularLinear":
100
+ masked = cls(
101
+ in_features=linear.in_features,
102
+ out_features=linear.out_features,
103
+ prefix_in_features=prefix_in_features,
104
+ prefix_out_features=prefix_out_features,
105
+ bias=linear.bias is not None,
106
+ device=linear.weight.device,
107
+ dtype=linear.weight.dtype,
108
+ )
109
+ with torch.no_grad():
110
+ masked.weight.copy_(linear.weight)
111
+ masked.weight.mul_(masked.weight_mask)
112
+ if linear.bias is not None:
113
+ masked.bias.copy_(linear.bias)
114
+ return masked
115
+
116
+ def reset_parameters(self) -> None:
117
+ nn.init.kaiming_uniform_(self.weight, a=5**0.5)
118
+ if self.bias is not None:
119
+ fan_in, _ = nn.init._calculate_fan_in_and_fan_out(self.weight)
120
+ bound = 1 / fan_in**0.5 if fan_in > 0 else 0
121
+ nn.init.uniform_(self.bias, -bound, bound)
122
+ with torch.no_grad():
123
+ self.weight.mul_(self.weight_mask)
124
+
125
+ def forward(self, input: torch.Tensor) -> torch.Tensor:
126
+ return F.linear(input, self.weight * self.weight_mask, self.bias)
127
+
128
+
129
+ class BlockPrefixNorm(nn.Module):
130
+ def __init__(self, hidden_size: int, prefix_hidden_size: int, eps: float = 1e-6):
131
+ super().__init__()
132
+ self.hidden_size = hidden_size
133
+ self.prefix_hidden_size = prefix_hidden_size
134
+ self.weight = nn.Parameter(torch.ones(hidden_size))
135
+ self.variance_epsilon = eps
136
+
137
+ @classmethod
138
+ def from_rms_norm(cls, norm: LlamaRMSNorm, prefix_hidden_size: int) -> "BlockPrefixNorm":
139
+ block_norm = cls(norm.weight.numel(), prefix_hidden_size, norm.variance_epsilon)
140
+ with torch.no_grad():
141
+ block_norm.weight.copy_(norm.weight)
142
+ return block_norm
143
+
144
+ def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
145
+ input_dtype = hidden_states.dtype
146
+ hidden_states = hidden_states.to(torch.float32)
147
+ prefix = hidden_states[..., : self.prefix_hidden_size]
148
+ suffix = hidden_states[..., self.prefix_hidden_size :]
149
+ prefix = prefix * torch.rsqrt(prefix.pow(2).mean(-1, keepdim=True) + self.variance_epsilon)
150
+ suffix = suffix * torch.rsqrt(hidden_states.pow(2).mean(-1, keepdim=True) + self.variance_epsilon)
151
+ return self.weight * torch.cat((prefix, suffix), dim=-1).to(input_dtype)
152
+
153
+
154
+ class BlockPrefixRMSNorm(BlockPrefixNorm):
155
+ """Backward-compatible alias; new models instantiate BlockPrefixNorm.
156
+
157
+ vLLM's Transformers backend replaces classes whose names end with
158
+ "RMSNorm". The prefix norm is behaviorally different from a standard RMSNorm,
159
+ so StairFormerModel uses BlockPrefixNorm to keep it intact under vLLM.
160
+ """
161
+
162
+
163
+ def apply_stairformer_modules(model: LlamaModel, config: LlamaConfig) -> None:
164
+ prefix_hidden_size = config.stairformer_prefix_hidden_size
165
+ prefix_intermediate_size = config.stairformer_prefix_intermediate_size
166
+ dense_entry_layers = config.stairformer_dense_entry_layers
167
+ prefix_key_value_size = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)
168
+
169
+ for layer_idx, layer in enumerate(model.layers):
170
+ if layer_idx < dense_entry_layers:
171
+ continue
172
+ layer.input_layernorm = BlockPrefixNorm.from_rms_norm(
173
+ layer.input_layernorm,
174
+ prefix_hidden_size,
175
+ )
176
+ layer.post_attention_layernorm = BlockPrefixNorm.from_rms_norm(
177
+ layer.post_attention_layernorm,
178
+ prefix_hidden_size,
179
+ )
180
+ attn = layer.self_attn
181
+ attn.q_proj = BlockLowerTriangularLinear.from_linear(attn.q_proj, prefix_hidden_size, prefix_hidden_size)
182
+ attn.k_proj = BlockLowerTriangularLinear.from_linear(attn.k_proj, prefix_hidden_size, prefix_key_value_size)
183
+ attn.v_proj = BlockLowerTriangularLinear.from_linear(attn.v_proj, prefix_hidden_size, prefix_key_value_size)
184
+ attn.o_proj = BlockLowerTriangularLinear.from_linear(attn.o_proj, prefix_hidden_size, prefix_hidden_size)
185
+ mlp = layer.mlp
186
+ mlp.gate_proj = BlockLowerTriangularLinear.from_linear(
187
+ mlp.gate_proj,
188
+ prefix_hidden_size,
189
+ prefix_intermediate_size,
190
+ )
191
+ mlp.up_proj = BlockLowerTriangularLinear.from_linear(
192
+ mlp.up_proj,
193
+ prefix_hidden_size,
194
+ prefix_intermediate_size,
195
+ )
196
+ mlp.down_proj = BlockLowerTriangularLinear.from_linear(
197
+ mlp.down_proj,
198
+ prefix_intermediate_size,
199
+ prefix_hidden_size,
200
+ )
201
+ model.norm = BlockPrefixNorm.from_rms_norm(model.norm, prefix_hidden_size)
202
+
203
+
204
+ class StairFormerModel(LlamaModel):
205
+ config_class = LlamaConfig
206
+ _supports_attention_backend = True
207
+ _keys_to_ignore_on_save = [r".*weight_mask"]
208
+ _keys_to_ignore_on_load_missing = [r".*weight_mask"]
209
+
210
+ def __init__(self, config: LlamaConfig):
211
+ configure_stairformer_config(config)
212
+ super().__init__(config)
213
+ apply_stairformer_modules(self, config)
214
+
215
+
216
+ class StairFormerForCausalLM(LlamaForCausalLM):
217
+ config_class = LlamaConfig
218
+ _supports_attention_backend = True
219
+ _keys_to_ignore_on_save = [r".*weight_mask"]
220
+ _keys_to_ignore_on_load_missing = [r".*weight_mask"]
221
+
222
+ def __init__(self, config: LlamaConfig):
223
+ configure_stairformer_config(config)
224
+ super().__init__(config)
225
+ self._apply_stairformer_modules()
226
+
227
+ @property
228
+ def prefix_hidden_size(self) -> int:
229
+ return self.config.stairformer_prefix_hidden_size
230
+
231
+ @property
232
+ def prefix_intermediate_size(self) -> int:
233
+ return self.config.stairformer_prefix_intermediate_size
234
+
235
+ @property
236
+ def dense_entry_layers(self) -> int:
237
+ return self.config.stairformer_dense_entry_layers
238
+
239
+ @property
240
+ def prefix_key_value_size(self) -> int:
241
+ return getattr(self.config, "head_dim", self.config.hidden_size // self.config.num_attention_heads)
242
+
243
+ @property
244
+ def nested_loss_alpha(self) -> float:
245
+ return float(self.config.stairformer_nested_loss_alpha)
246
+
247
+ def _apply_stairformer_modules(self) -> None:
248
+ apply_stairformer_modules(self.model, self.config)
249
+
250
+ def nested_lm_head_weight(self) -> torch.Tensor:
251
+ return self.lm_head.weight[:, : self.prefix_hidden_size]
252
+
253
+ def nested_logits_from_hidden(self, hidden_states: torch.Tensor, num_logits_to_keep: int = 0) -> torch.Tensor:
254
+ prefix_hidden = hidden_states[:, -num_logits_to_keep:, : self.prefix_hidden_size]
255
+ return F.linear(prefix_hidden, self.nested_lm_head_weight()).float()
256
+
257
+ def forward(
258
+ self,
259
+ input_ids: torch.LongTensor = None,
260
+ attention_mask: Optional[torch.Tensor] = None,
261
+ position_ids: Optional[torch.LongTensor] = None,
262
+ past_key_values: Optional[Union[List[torch.FloatTensor], Tuple[Tuple[torch.FloatTensor]]]] = None,
263
+ inputs_embeds: Optional[torch.FloatTensor] = None,
264
+ labels: Optional[torch.LongTensor] = None,
265
+ use_cache: Optional[bool] = None,
266
+ output_attentions: Optional[bool] = None,
267
+ output_hidden_states: Optional[bool] = None,
268
+ return_dict: Optional[bool] = None,
269
+ cache_position: Optional[torch.LongTensor] = None,
270
+ num_logits_to_keep: int = 0,
271
+ output_nested_logits: bool = False,
272
+ ):
273
+ outputs = self.model(
274
+ input_ids=input_ids,
275
+ attention_mask=attention_mask,
276
+ position_ids=position_ids,
277
+ past_key_values=past_key_values,
278
+ inputs_embeds=inputs_embeds,
279
+ use_cache=use_cache,
280
+ output_attentions=output_attentions,
281
+ output_hidden_states=output_hidden_states,
282
+ return_dict=True,
283
+ cache_position=cache_position,
284
+ )
285
+ hidden_states = outputs[0]
286
+ logits = self.lm_head(hidden_states[:, -num_logits_to_keep:, :]).float()
287
+ nested_logits = self.nested_logits_from_hidden(hidden_states, num_logits_to_keep) if (output_nested_logits or labels is not None) else None
288
+ loss = loss_full = loss_nested = None
289
+ if labels is not None:
290
+ loss_fct = CrossEntropyLoss()
291
+ shift_labels = labels[..., 1:].contiguous().view(-1)
292
+ loss_full = loss_fct(logits[..., :-1, :].contiguous().view(-1, self.config.vocab_size), shift_labels.to(logits.device))
293
+ loss_nested = loss_fct(nested_logits[..., :-1, :].contiguous().view(-1, self.config.vocab_size), shift_labels.to(nested_logits.device))
294
+ alpha = self.nested_loss_alpha
295
+ loss = (1.0 - alpha) * loss_full + alpha * loss_nested
296
+ return StairFormerCausalLMOutputWithPast(
297
+ loss=loss,
298
+ loss_full=loss_full,
299
+ loss_nested=loss_nested,
300
+ logits=logits,
301
+ nested_logits=nested_logits if output_nested_logits else None,
302
+ past_key_values=outputs.past_key_values,
303
+ hidden_states=outputs.hidden_states,
304
+ attentions=outputs.attentions,
305
+ )
306
+
307
+
308
+ class AsymmetricNestedLlamaModel(LlamaPreTrainedModel):
309
+ config_class = LlamaConfig
310
+ _supports_attention_backend = True
311
+
312
+ def __init__(self, config: LlamaConfig):
313
+ super().__init__(config)
314
+ configure_stairformer_config(config)
315
+ self.prefix_hidden_size = config.stairformer_prefix_hidden_size
316
+ self.prefix_intermediate_size = config.stairformer_prefix_intermediate_size
317
+ self.dense_entry_layers = config.stairformer_dense_entry_layers
318
+ self.vocab_size = config.vocab_size
319
+ self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, config.pad_token_id)
320
+ self.entry_rotary_emb = LlamaRotaryEmbedding(config)
321
+ self.entry_layers = nn.ModuleList(
322
+ [LlamaDecoderLayer(config, layer_idx=i) for i in range(self.dense_entry_layers)]
323
+ )
324
+ self.prefix_config = self._make_prefix_config(config)
325
+ self.prefix_rotary_emb = LlamaRotaryEmbedding(self.prefix_config)
326
+ self.layers = nn.ModuleList(
327
+ [
328
+ LlamaDecoderLayer(self.prefix_config, layer_idx=self.dense_entry_layers + i)
329
+ for i in range(config.num_hidden_layers - self.dense_entry_layers)
330
+ ]
331
+ )
332
+ self.norm = LlamaRMSNorm(self.prefix_hidden_size, eps=config.rms_norm_eps)
333
+ self.post_init()
334
+
335
+ def get_input_embeddings(self):
336
+ return self.embed_tokens
337
+
338
+ def set_input_embeddings(self, value):
339
+ self.embed_tokens = value
340
+
341
+ def get_decoder(self):
342
+ return self
343
+
344
+ def _make_prefix_config(self, config: LlamaConfig) -> LlamaConfig:
345
+ prefix_config = LlamaConfig(**config.to_dict())
346
+ head_dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)
347
+ prefix_config.hidden_size = self.prefix_hidden_size
348
+ prefix_config.intermediate_size = self.prefix_intermediate_size
349
+ prefix_config.num_attention_heads = self.prefix_hidden_size // head_dim
350
+ prefix_config.num_key_value_heads = 1
351
+ prefix_config.tie_word_embeddings = False
352
+ if getattr(prefix_config, "_attn_implementation", None) is None:
353
+ prefix_config._attn_implementation = "eager"
354
+ return prefix_config
355
+
356
+ @staticmethod
357
+ def _make_causal_mask(
358
+ batch_size: int,
359
+ seq_len: int,
360
+ dtype: torch.dtype,
361
+ device: torch.device,
362
+ attention_mask: Optional[torch.Tensor] = None,
363
+ ):
364
+ min_value = torch.finfo(dtype).min
365
+ mask = torch.full((seq_len, seq_len), min_value, dtype=dtype, device=device)
366
+ mask = torch.triu(mask, diagonal=1)
367
+ mask = mask[None, None, :, :].expand(batch_size, 1, seq_len, seq_len)
368
+ if attention_mask is not None:
369
+ padding_mask = (1.0 - attention_mask[:, None, None, :].to(dtype)) * min_value
370
+ mask = mask + padding_mask
371
+ return mask
372
+
373
+ @staticmethod
374
+ def _run_decoder_layer(layer: nn.Module, hidden_states: torch.Tensor, **kwargs) -> torch.Tensor:
375
+ try:
376
+ layer_outputs = layer(hidden_states, **kwargs)
377
+ except TypeError:
378
+ kwargs.pop("position_embeddings", None)
379
+ layer_outputs = layer(hidden_states, **kwargs)
380
+ return layer_outputs[0] if isinstance(layer_outputs, tuple) else layer_outputs
381
+
382
+ def forward(
383
+ self,
384
+ input_ids: torch.LongTensor = None,
385
+ attention_mask: Optional[torch.Tensor] = None,
386
+ position_ids: Optional[torch.LongTensor] = None,
387
+ past_key_values: Optional[Union[List[torch.FloatTensor], Tuple[Tuple[torch.FloatTensor]]]] = None,
388
+ inputs_embeds: Optional[torch.FloatTensor] = None,
389
+ labels: Optional[torch.LongTensor] = None,
390
+ use_cache: Optional[bool] = None,
391
+ output_attentions: Optional[bool] = None,
392
+ output_hidden_states: Optional[bool] = None,
393
+ return_dict: bool = True,
394
+ cache_position: Optional[torch.LongTensor] = None,
395
+ num_logits_to_keep: int = 0,
396
+ **kwargs,
397
+ ):
398
+ del output_attentions, output_hidden_states, cache_position, num_logits_to_keep
399
+ if past_key_values is not None:
400
+ raise NotImplementedError("AsymmetricNestedLlamaModel does not support KV cache reuse yet.")
401
+ if (input_ids is None) == (inputs_embeds is None):
402
+ raise ValueError("Specify exactly one of input_ids or inputs_embeds.")
403
+ if inputs_embeds is None:
404
+ hidden_states = self.embed_tokens(input_ids)
405
+ else:
406
+ hidden_states = inputs_embeds
407
+ batch_size, seq_len, _ = hidden_states.shape
408
+ if position_ids is None:
409
+ position_ids = torch.arange(seq_len, device=hidden_states.device).unsqueeze(0)
410
+ causal_mask = self._make_causal_mask(
411
+ batch_size,
412
+ seq_len,
413
+ hidden_states.dtype,
414
+ hidden_states.device,
415
+ attention_mask,
416
+ )
417
+ entry_position_embeddings = self.entry_rotary_emb(hidden_states, position_ids)
418
+ for layer in self.entry_layers:
419
+ hidden_states = self._run_decoder_layer(
420
+ layer,
421
+ hidden_states,
422
+ attention_mask=causal_mask,
423
+ position_ids=position_ids,
424
+ use_cache=False,
425
+ position_embeddings=entry_position_embeddings,
426
+ **kwargs,
427
+ )
428
+ hidden_states = hidden_states[..., : self.prefix_hidden_size]
429
+ prefix_position_embeddings = self.prefix_rotary_emb(hidden_states, position_ids)
430
+ for layer in self.layers:
431
+ hidden_states = self._run_decoder_layer(
432
+ layer,
433
+ hidden_states,
434
+ attention_mask=causal_mask,
435
+ position_ids=position_ids,
436
+ use_cache=False,
437
+ position_embeddings=prefix_position_embeddings,
438
+ **kwargs,
439
+ )
440
+ hidden_states = self.norm(hidden_states)
441
+ if not return_dict:
442
+ return (hidden_states,)
443
+ return BaseModelOutputWithPast(
444
+ last_hidden_state=hidden_states,
445
+ past_key_values=None if use_cache else None,
446
+ )
447
+
448
+
449
+ class AsymmetricNestedLlamaForCausalLM(LlamaPreTrainedModel, GenerationMixin):
450
+ config_class = LlamaConfig
451
+ _supports_attention_backend = True
452
+
453
+ def __init__(self, config: LlamaConfig):
454
+ super().__init__(config)
455
+ configure_stairformer_config(config)
456
+ self.prefix_hidden_size = config.stairformer_prefix_hidden_size
457
+ self.prefix_intermediate_size = config.stairformer_prefix_intermediate_size
458
+ self.dense_entry_layers = config.stairformer_dense_entry_layers
459
+ self.vocab_size = config.vocab_size
460
+ self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, config.pad_token_id)
461
+ self.entry_rotary_emb = LlamaRotaryEmbedding(config)
462
+ self.entry_layers = nn.ModuleList(
463
+ [LlamaDecoderLayer(config, layer_idx=i) for i in range(self.dense_entry_layers)]
464
+ )
465
+ self.prefix_config = self._make_prefix_config(config)
466
+ self.prefix_rotary_emb = LlamaRotaryEmbedding(self.prefix_config)
467
+ self.layers = nn.ModuleList(
468
+ [
469
+ LlamaDecoderLayer(self.prefix_config, layer_idx=self.dense_entry_layers + i)
470
+ for i in range(config.num_hidden_layers - self.dense_entry_layers)
471
+ ]
472
+ )
473
+ self.norm = LlamaRMSNorm(self.prefix_hidden_size, eps=config.rms_norm_eps)
474
+ self.lm_head = nn.Linear(self.prefix_hidden_size, self.vocab_size, bias=False)
475
+ self.post_init()
476
+
477
+ get_input_embeddings = AsymmetricNestedLlamaModel.get_input_embeddings
478
+ set_input_embeddings = AsymmetricNestedLlamaModel.set_input_embeddings
479
+ _make_prefix_config = AsymmetricNestedLlamaModel._make_prefix_config
480
+ _make_causal_mask = staticmethod(AsymmetricNestedLlamaModel._make_causal_mask)
481
+ _run_decoder_layer = staticmethod(AsymmetricNestedLlamaModel._run_decoder_layer)
482
+
483
+ def get_output_embeddings(self):
484
+ return self.lm_head
485
+
486
+ def set_output_embeddings(self, new_embeddings):
487
+ self.lm_head = new_embeddings
488
+
489
+ def forward(
490
+ self,
491
+ input_ids: torch.LongTensor = None,
492
+ attention_mask: Optional[torch.Tensor] = None,
493
+ position_ids: Optional[torch.LongTensor] = None,
494
+ past_key_values: Optional[Union[List[torch.FloatTensor], Tuple[Tuple[torch.FloatTensor]]]] = None,
495
+ inputs_embeds: Optional[torch.FloatTensor] = None,
496
+ labels: Optional[torch.LongTensor] = None,
497
+ use_cache: Optional[bool] = None,
498
+ output_attentions: Optional[bool] = None,
499
+ output_hidden_states: Optional[bool] = None,
500
+ return_dict: bool = True,
501
+ cache_position: Optional[torch.LongTensor] = None,
502
+ num_logits_to_keep: int = 0,
503
+ **kwargs,
504
+ ):
505
+ del output_attentions, output_hidden_states, cache_position
506
+ if past_key_values is not None:
507
+ raise NotImplementedError("AsymmetricNestedLlamaForCausalLM does not support KV cache reuse yet.")
508
+ if (input_ids is None) == (inputs_embeds is None):
509
+ raise ValueError("Specify exactly one of input_ids or inputs_embeds.")
510
+ if inputs_embeds is None:
511
+ hidden_states = self.embed_tokens(input_ids)
512
+ else:
513
+ hidden_states = inputs_embeds
514
+ batch_size, seq_len, _ = hidden_states.shape
515
+ if position_ids is None:
516
+ position_ids = torch.arange(seq_len, device=hidden_states.device).unsqueeze(0)
517
+ causal_mask = self._make_causal_mask(
518
+ batch_size,
519
+ seq_len,
520
+ hidden_states.dtype,
521
+ hidden_states.device,
522
+ attention_mask,
523
+ )
524
+ entry_position_embeddings = self.entry_rotary_emb(hidden_states, position_ids)
525
+ for layer in self.entry_layers:
526
+ hidden_states = self._run_decoder_layer(
527
+ layer,
528
+ hidden_states,
529
+ attention_mask=causal_mask,
530
+ position_ids=position_ids,
531
+ use_cache=False,
532
+ position_embeddings=entry_position_embeddings,
533
+ **kwargs,
534
+ )
535
+ hidden_states = hidden_states[..., : self.prefix_hidden_size]
536
+ prefix_position_embeddings = self.prefix_rotary_emb(hidden_states, position_ids)
537
+ for layer in self.layers:
538
+ hidden_states = self._run_decoder_layer(
539
+ layer,
540
+ hidden_states,
541
+ attention_mask=causal_mask,
542
+ position_ids=position_ids,
543
+ use_cache=False,
544
+ position_embeddings=prefix_position_embeddings,
545
+ **kwargs,
546
+ )
547
+ hidden_states = self.norm(hidden_states)
548
+ logits = self.lm_head(hidden_states[:, -num_logits_to_keep:, :]).float()
549
+ loss = None
550
+ if labels is not None:
551
+ loss = CrossEntropyLoss()(
552
+ logits[..., :-1, :].contiguous().view(-1, self.vocab_size),
553
+ labels[..., 1:].contiguous().view(-1).to(logits.device),
554
+ )
555
+ if not return_dict:
556
+ output = (logits, hidden_states)
557
+ return (loss,) + output if loss is not None else output
558
+ return AsymmetricNestedLlamaOutput(
559
+ loss=loss,
560
+ logits=logits,
561
+ hidden_states=hidden_states,
562
+ past_key_values=None if use_cache else None,
563
+ )
564
+
565
+ def prepare_inputs_for_generation(
566
+ self,
567
+ input_ids,
568
+ past_key_values=None,
569
+ attention_mask=None,
570
+ inputs_embeds=None,
571
+ **kwargs,
572
+ ):
573
+ del past_key_values
574
+ model_inputs = {"inputs_embeds": inputs_embeds} if inputs_embeds is not None else {"input_ids": input_ids}
575
+ model_inputs.update(
576
+ attention_mask=attention_mask,
577
+ use_cache=False,
578
+ )
579
+ return model_inputs
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer_config.json ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": false,
3
+ "backend": "tokenizers",
4
+ "bos_token": "<|endoftext|>",
5
+ "eos_token": "<|endoftext|>",
6
+ "errors": "replace",
7
+ "is_local": false,
8
+ "local_files_only": false,
9
+ "model_max_length": 2048,
10
+ "pad_token": "<|endoftext|>",
11
+ "tokenizer_class": "GPT2Tokenizer",
12
+ "unk_token": "<|endoftext|>"
13
+ }