Arsh9210 commited on
Commit
c7cc693
·
verified ·
1 Parent(s): 0093dca

Added nemotron_dense_vllm_plugin/nemotron_dense_vllm/nemotron_dense.py

Browse files
nemotron_dense_vllm_plugin/nemotron_dense_vllm/nemotron_dense.py ADDED
@@ -0,0 +1,423 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # SPDX-License-Identifier: Apache-2.0
2
+ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project
3
+
4
+ from collections.abc import Iterable
5
+ from itertools import islice
6
+
7
+ import torch
8
+ from torch import nn
9
+
10
+ from vllm.model_executor.layers.attention import Attention
11
+ from vllm.compilation.decorators import support_torch_compile
12
+ from vllm.config import CacheConfig, VllmConfig
13
+ from vllm.distributed import get_pp_group, get_tensor_model_parallel_world_size
14
+ from vllm.model_executor.layers.activation import get_act_fn
15
+ from vllm.model_executor.layers.linear import (
16
+ ColumnParallelLinear,
17
+ QKVParallelLinear,
18
+ RowParallelLinear,
19
+ )
20
+ from vllm.model_executor.layers.logits_processor import LogitsProcessor
21
+ from vllm.model_executor.layers.quantization import QuantizationConfig
22
+ from vllm.model_executor.layers.rotary_embedding import get_rope
23
+ from vllm.model_executor.layers.vocab_parallel_embedding import (
24
+ ParallelLMHead,
25
+ VocabParallelEmbedding,
26
+ )
27
+ from vllm.model_executor.model_loader.weight_utils import (
28
+ default_weight_loader,
29
+ maybe_remap_kv_scale_name,
30
+ )
31
+ from vllm.sequence import IntermediateTensors
32
+ from vllm.transformers_utils.configs.nemotron import NemotronConfig
33
+
34
+ from vllm.model_executor.models.interfaces import SupportsLoRA, SupportsPP
35
+ from vllm.model_executor.models.utils import (
36
+ AutoWeightsLoader,
37
+ PPMissingLayer,
38
+ is_pp_missing_parameter,
39
+ make_empty_intermediate_tensors_factory,
40
+ make_layers,
41
+ maybe_prefix,
42
+ )
43
+
44
+
45
+ class NemotronDenseRMSNorm(nn.Module):
46
+ def __init__(self, hidden_size, eps=1e-5):
47
+ super().__init__()
48
+ self.weight = nn.Parameter(torch.ones(hidden_size))
49
+ self.variance_epsilon = eps
50
+
51
+ def forward(self, hidden_states):
52
+ input_dtype = hidden_states.dtype
53
+ hidden_states = hidden_states.to(torch.float32)
54
+ variance = hidden_states.pow(2).mean(-1, keepdim=True)
55
+ hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon)
56
+ return self.weight * hidden_states.to(input_dtype)
57
+
58
+
59
+ class NemotronDenseMLP(nn.Module):
60
+ def __init__(
61
+ self,
62
+ hidden_size: int,
63
+ intermediate_size: int,
64
+ hidden_act: str,
65
+ quant_config: QuantizationConfig | None = None,
66
+ bias: bool = False,
67
+ prefix: str = "",
68
+ ) -> None:
69
+ super().__init__()
70
+ self.up_proj = ColumnParallelLinear(
71
+ input_size=hidden_size,
72
+ output_size=intermediate_size,
73
+ bias=bias,
74
+ quant_config=quant_config,
75
+ prefix=f"{prefix}.up_proj",
76
+ )
77
+ self.down_proj = RowParallelLinear(
78
+ input_size=intermediate_size,
79
+ output_size=hidden_size,
80
+ bias=bias,
81
+ quant_config=quant_config,
82
+ prefix=f"{prefix}.down_proj",
83
+ )
84
+ self.act_fn = get_act_fn(hidden_act)
85
+
86
+ def forward(self, x):
87
+ up, _ = self.up_proj(x)
88
+ x = self.act_fn(up)
89
+ x, _ = self.down_proj(x)
90
+ return x
91
+
92
+
93
+ class NemotronDenseAttention(nn.Module):
94
+ def __init__(
95
+ self,
96
+ config: NemotronConfig,
97
+ hidden_size: int,
98
+ num_heads: int,
99
+ num_kv_heads: int,
100
+ max_position_embeddings: int = 8192,
101
+ quant_config: QuantizationConfig | None = None,
102
+ bias: bool = False,
103
+ cache_config: CacheConfig | None = None,
104
+ prefix: str = "",
105
+ ) -> None:
106
+ super().__init__()
107
+ self.hidden_size = hidden_size
108
+ tp_size = get_tensor_model_parallel_world_size()
109
+ self.total_num_heads = num_heads
110
+ assert self.total_num_heads % tp_size == 0
111
+ self.num_heads = self.total_num_heads // tp_size
112
+ self.total_num_kv_heads = num_kv_heads
113
+ if self.total_num_kv_heads >= tp_size:
114
+ assert self.total_num_kv_heads % tp_size == 0
115
+ else:
116
+ assert tp_size % self.total_num_kv_heads == 0
117
+ self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size)
118
+ self.head_dim = getattr(config, "head_dim", None)
119
+ if self.head_dim is None:
120
+ self.head_dim = self.hidden_size // self.total_num_heads
121
+ self.q_size = self.num_heads * self.head_dim
122
+ self.kv_size = self.num_kv_heads * self.head_dim
123
+ self.scaling = self.head_dim**-0.5
124
+ self.max_position_embeddings = max_position_embeddings
125
+
126
+ self.qkv_proj = QKVParallelLinear(
127
+ hidden_size=hidden_size,
128
+ head_size=self.head_dim,
129
+ total_num_heads=self.total_num_heads,
130
+ total_num_kv_heads=self.total_num_kv_heads,
131
+ bias=bias,
132
+ quant_config=quant_config,
133
+ prefix=f"{prefix}.qkv_proj",
134
+ )
135
+ self.o_proj = RowParallelLinear(
136
+ input_size=self.total_num_heads * self.head_dim,
137
+ output_size=hidden_size,
138
+ bias=bias,
139
+ quant_config=quant_config,
140
+ prefix=f"{prefix}.o_proj",
141
+ )
142
+
143
+ self.rotary_emb = get_rope(
144
+ self.head_dim,
145
+ max_position=max_position_embeddings,
146
+ rope_parameters=config.rope_parameters,
147
+ )
148
+ self.attn = Attention(
149
+ self.num_heads,
150
+ self.head_dim,
151
+ self.scaling,
152
+ num_kv_heads=self.num_kv_heads,
153
+ cache_config=cache_config,
154
+ quant_config=quant_config,
155
+ prefix=f"{prefix}.attn",
156
+ )
157
+
158
+ def forward(
159
+ self,
160
+ positions: torch.Tensor,
161
+ hidden_states: torch.Tensor,
162
+ ) -> torch.Tensor:
163
+ qkv, _ = self.qkv_proj(hidden_states)
164
+ q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
165
+ q, k = self.rotary_emb(positions, q, k)
166
+ attn_output = self.attn(q, k, v)
167
+ output, _ = self.o_proj(attn_output)
168
+ return output
169
+
170
+
171
+ class NemotronDenseDecoderLayer(nn.Module):
172
+ def __init__(
173
+ self,
174
+ config: NemotronConfig,
175
+ cache_config: CacheConfig | None = None,
176
+ quant_config: QuantizationConfig | None = None,
177
+ prefix: str = "",
178
+ ) -> None:
179
+ super().__init__()
180
+ self.hidden_size = config.hidden_size
181
+ max_position_embeddings = getattr(config, "max_position_embeddings", 8192)
182
+ attention_bias = getattr(config, "attention_bias", False) or getattr(
183
+ config, "bias", False
184
+ )
185
+ self.self_attn = NemotronDenseAttention(
186
+ config=config,
187
+ hidden_size=self.hidden_size,
188
+ num_heads=config.num_attention_heads,
189
+ num_kv_heads=getattr(
190
+ config, "num_key_value_heads", config.num_attention_heads
191
+ ),
192
+ max_position_embeddings=max_position_embeddings,
193
+ quant_config=quant_config,
194
+ bias=attention_bias,
195
+ cache_config=cache_config,
196
+ prefix=f"{prefix}.self_attn",
197
+ )
198
+ self.mlp = NemotronDenseMLP(
199
+ hidden_size=self.hidden_size,
200
+ intermediate_size=config.intermediate_size,
201
+ hidden_act=config.hidden_act,
202
+ quant_config=quant_config,
203
+ bias=getattr(config, "mlp_bias", False),
204
+ prefix=f"{prefix}.mlp",
205
+ )
206
+ self.input_layernorm = NemotronDenseRMSNorm(
207
+ config.hidden_size, eps=config.norm_eps
208
+ )
209
+ self.post_attention_layernorm = NemotronDenseRMSNorm(
210
+ config.hidden_size, eps=config.norm_eps
211
+ )
212
+
213
+ def forward(
214
+ self,
215
+ positions: torch.Tensor,
216
+ hidden_states: torch.Tensor,
217
+ ) -> torch.Tensor:
218
+ residual = hidden_states
219
+ hidden_states = self.input_layernorm(hidden_states)
220
+ hidden_states = self.self_attn(
221
+ positions=positions,
222
+ hidden_states=hidden_states,
223
+ )
224
+ hidden_states = residual + hidden_states
225
+
226
+ residual = hidden_states
227
+ hidden_states = self.post_attention_layernorm(hidden_states)
228
+ hidden_states = self.mlp(hidden_states)
229
+ hidden_states = residual + hidden_states
230
+ return hidden_states
231
+
232
+
233
+ @support_torch_compile
234
+ class NemotronDenseModel(nn.Module):
235
+ def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
236
+ super().__init__()
237
+
238
+ config = vllm_config.model_config.hf_config
239
+ cache_config = vllm_config.cache_config
240
+ quant_config = vllm_config.quant_config
241
+
242
+ self.config = config
243
+ self.quant_config = quant_config
244
+
245
+ self.vocab_size = config.vocab_size
246
+
247
+ if get_pp_group().is_first_rank or (
248
+ config.tie_word_embeddings and get_pp_group().is_last_rank
249
+ ):
250
+ self.embed_tokens = VocabParallelEmbedding(
251
+ self.vocab_size,
252
+ config.hidden_size,
253
+ )
254
+ else:
255
+ self.embed_tokens = PPMissingLayer()
256
+ self.start_layer, self.end_layer, self.layers = make_layers(
257
+ config.num_hidden_layers,
258
+ lambda prefix: NemotronDenseDecoderLayer(
259
+ config=config,
260
+ cache_config=cache_config,
261
+ quant_config=quant_config,
262
+ prefix=prefix,
263
+ ),
264
+ prefix=f"{prefix}.layers",
265
+ )
266
+ if get_pp_group().is_last_rank:
267
+ self.norm = NemotronDenseRMSNorm(config.hidden_size, eps=config.norm_eps)
268
+ else:
269
+ self.norm = PPMissingLayer()
270
+
271
+ self.make_empty_intermediate_tensors = make_empty_intermediate_tensors_factory(
272
+ ["hidden_states"], config.hidden_size
273
+ )
274
+
275
+ def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
276
+ return self.embed_tokens(input_ids)
277
+
278
+ def forward(
279
+ self,
280
+ input_ids: torch.Tensor | None,
281
+ positions: torch.Tensor,
282
+ intermediate_tensors: IntermediateTensors | None,
283
+ inputs_embeds: torch.Tensor | None = None,
284
+ ) -> torch.Tensor | IntermediateTensors:
285
+ if get_pp_group().is_first_rank:
286
+ if inputs_embeds is not None:
287
+ hidden_states = inputs_embeds
288
+ else:
289
+ hidden_states = self.embed_input_ids(input_ids)
290
+ else:
291
+ assert intermediate_tensors is not None
292
+ hidden_states = intermediate_tensors["hidden_states"]
293
+
294
+ for layer in islice(self.layers, self.start_layer, self.end_layer):
295
+ hidden_states = layer(positions, hidden_states)
296
+
297
+ if not get_pp_group().is_last_rank:
298
+ return IntermediateTensors({"hidden_states": hidden_states})
299
+
300
+ hidden_states = self.norm(hidden_states)
301
+ return hidden_states
302
+
303
+ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
304
+ stacked_params_mapping = [
305
+ (".qkv_proj", ".q_proj", "q"),
306
+ (".qkv_proj", ".k_proj", "k"),
307
+ (".qkv_proj", ".v_proj", "v"),
308
+ ]
309
+ params_dict = dict(self.named_parameters())
310
+ loaded_params: set[str] = set()
311
+ for name, loaded_weight in weights:
312
+ if self.quant_config is not None and (
313
+ scale_name := self.quant_config.get_cache_scale(name)
314
+ ):
315
+ param = params_dict[scale_name]
316
+ weight_loader = getattr(param, "weight_loader", default_weight_loader)
317
+ loaded_weight = (
318
+ loaded_weight if loaded_weight.dim() == 0 else loaded_weight[0]
319
+ )
320
+ weight_loader(param, loaded_weight)
321
+ loaded_params.add(scale_name)
322
+ continue
323
+ for param_name, weight_name, shard_id in stacked_params_mapping:
324
+ if weight_name not in name:
325
+ continue
326
+ name = name.replace(weight_name, param_name)
327
+ if name.endswith(".bias") and name not in params_dict:
328
+ continue
329
+
330
+ if is_pp_missing_parameter(name, self):
331
+ continue
332
+
333
+ param = params_dict[name]
334
+ weight_loader = param.weight_loader
335
+ weight_loader(param, loaded_weight, shard_id)
336
+
337
+ break
338
+ else:
339
+ if name.endswith(".bias") and name not in params_dict:
340
+ continue
341
+ name = maybe_remap_kv_scale_name(name, params_dict)
342
+ if name is None:
343
+ continue
344
+
345
+ if is_pp_missing_parameter(name, self):
346
+ continue
347
+
348
+ param = params_dict[name]
349
+ weight_loader = getattr(param, "weight_loader", default_weight_loader)
350
+ weight_loader(param, loaded_weight)
351
+ loaded_params.add(name)
352
+ return loaded_params
353
+
354
+
355
+ class NemotronDenseForCausalLM(nn.Module, SupportsLoRA, SupportsPP):
356
+ packed_modules_mapping = {
357
+ "qkv_proj": [
358
+ "q_proj",
359
+ "k_proj",
360
+ "v_proj",
361
+ ],
362
+ }
363
+
364
+ embedding_modules = {
365
+ "embed_tokens": "input_embeddings",
366
+ "lm_head": "output_embeddings",
367
+ }
368
+
369
+ def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
370
+ super().__init__()
371
+ config = vllm_config.model_config.hf_config
372
+ quant_config = vllm_config.quant_config
373
+
374
+ assert hasattr(config, "norm_eps") and hasattr(config, "num_hidden_layers")
375
+
376
+ self.config = config
377
+ self.quant_config = quant_config
378
+
379
+ self.model = NemotronDenseModel(
380
+ vllm_config=vllm_config, prefix=maybe_prefix(prefix, "model")
381
+ )
382
+ if get_pp_group().is_last_rank:
383
+ self.lm_head = ParallelLMHead(
384
+ config.vocab_size,
385
+ config.hidden_size,
386
+ quant_config=quant_config,
387
+ prefix=maybe_prefix(prefix, "lm_head"),
388
+ )
389
+ if config.tie_word_embeddings:
390
+ self.lm_head.weight = self.model.embed_tokens.weight
391
+
392
+ logit_scale = getattr(config, "logit_scale", 1.0)
393
+ self.logits_processor = LogitsProcessor(
394
+ config.vocab_size, scale=logit_scale
395
+ )
396
+ else:
397
+ self.lm_head = PPMissingLayer()
398
+
399
+ self.make_empty_intermediate_tensors = (
400
+ self.model.make_empty_intermediate_tensors
401
+ )
402
+
403
+ def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor:
404
+ return self.model.embed_input_ids(input_ids)
405
+
406
+ def forward(
407
+ self,
408
+ input_ids: torch.Tensor,
409
+ positions: torch.Tensor,
410
+ intermediate_tensors: IntermediateTensors | None = None,
411
+ inputs_embeds: torch.Tensor | None = None,
412
+ ) -> torch.Tensor | IntermediateTensors:
413
+ return self.model(input_ids, positions, intermediate_tensors, inputs_embeds)
414
+
415
+ def compute_logits(
416
+ self,
417
+ hidden_states: torch.Tensor,
418
+ ) -> torch.Tensor | None:
419
+ return self.logits_processor(self.lm_head, hidden_states)
420
+
421
+ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
422
+ loader = AutoWeightsLoader(self)
423
+ return loader.load_weights(weights)