AlexHung29629 commited on
Commit
ea056a6
·
verified ·
1 Parent(s): a490733

Create modeling_pica.py

Browse files
Files changed (1) hide show
  1. modeling_pica.py +477 -0
modeling_pica.py ADDED
@@ -0,0 +1,477 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Callable, Optional, Tuple, Unpack, Union
2
+ import torch
3
+ from torch import nn
4
+ from transformers import AutoConfig, AutoModel, AutoModelForCausalLM
5
+ from transformers.utils import logging, LossKwargs
6
+ from transformers.cache_utils import Cache, DynamicCache, StaticCache
7
+ from transformers.models.llama.modeling_llama import LlamaRMSNorm, LlamaModel, LlamaMLP
8
+ from transformers.modeling_utils import PreTrainedModel
9
+ from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast
10
+ from transformers.generation import GenerationMixin
11
+ try:
12
+ from flash_sigmoid import flash_attn_func as flash_sigmoid_func
13
+ except:
14
+ flash_sigmoid_func = None
15
+ from .configuration_pica import PicaConfig
16
+
17
+ logger = logging.get_logger(__name__)
18
+
19
+
20
+ class PicaAttention(nn.Module):
21
+ """Multi-headed attention from 'Attention Is All You Need' paper"""
22
+
23
+ def __init__(self, config: PicaConfig, layer_idx: int):
24
+ super().__init__()
25
+ self.config = config
26
+ self.layer_idx = layer_idx
27
+ self.head_dim = getattr(config, "head_dim", config.hidden_size // config.num_attention_heads)
28
+ self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads
29
+ self.scaling = self.head_dim**-0.5
30
+ self.attention_dropout = config.attention_dropout
31
+ self.is_causal = True
32
+
33
+ self.q_proj = nn.Linear(
34
+ config.hidden_size, config.num_attention_heads * self.head_dim, bias=config.attention_bias
35
+ )
36
+ self.k_proj = nn.Linear(
37
+ config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias
38
+ )
39
+ self.v_proj = nn.Linear(
40
+ config.hidden_size, config.num_key_value_heads * self.head_dim, bias=config.attention_bias
41
+ )
42
+ self.o_proj = nn.Linear(
43
+ config.num_attention_heads * self.head_dim, config.hidden_size, bias=config.attention_bias
44
+ )
45
+
46
+ def forward(
47
+ self,
48
+ hidden_states: torch.Tensor,
49
+ attention_mask: Optional[torch.Tensor],
50
+ past_key_value: Optional[Cache] = None,
51
+ cache_position: Optional[torch.LongTensor] = None,
52
+ **kwargs,
53
+ ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]:
54
+ input_shape = hidden_states.shape[:-1]
55
+ hidden_shape = (*input_shape, -1, self.head_dim)
56
+
57
+ query_states = self.q_proj(hidden_states).view(hidden_shape).transpose(1, 2)
58
+ key_states = self.k_proj(hidden_states).view(hidden_shape).transpose(1, 2)
59
+ value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)
60
+
61
+ if past_key_value is not None:
62
+ # cache_position needed for the static cache
63
+ cache_kwargs = {"cache_position": cache_position}
64
+ key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs)
65
+
66
+ attention_interface: Callable = sigmoid_attention_forward
67
+
68
+ attn_output, attn_weights = attention_interface(
69
+ self,
70
+ query_states,
71
+ key_states,
72
+ value_states,
73
+ attention_mask,
74
+ dropout=0.0 if not self.training else self.attention_dropout,
75
+ scaling=self.scaling,
76
+ **kwargs,
77
+ )
78
+
79
+ attn_output = attn_output.reshape(*input_shape, -1).contiguous()
80
+ attn_output = self.o_proj(attn_output)
81
+ return attn_output, attn_weights
82
+
83
+ def sigmoid_attention_forward(
84
+ module: nn.Module,
85
+ query: torch.Tensor,
86
+ key: torch.Tensor,
87
+ value: torch.Tensor,
88
+ attention_mask: Optional[torch.Tensor],
89
+ scaling: float,
90
+ dropout: float = 0.0,
91
+ **kwargs,
92
+ ):
93
+ key_states = repeat_kv(key, module.num_key_value_groups)
94
+ value_states = repeat_kv(value, module.num_key_value_groups)
95
+
96
+ if flash_sigmoid_func is None:
97
+ attn_weights = torch.matmul(query, key_states.transpose(2, 3)) * scaling
98
+ if attention_mask is not None:
99
+ causal_mask = attention_mask[:, :, :, : key_states.shape[-2]]
100
+ attn_weights = attn_weights + causal_mask
101
+
102
+ attn_weights = nn.functional.sigmoid(attn_weights - torch.log(attn_weights.size(-2))).to(query.dtype)
103
+ attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training)
104
+ attn_output = torch.matmul(attn_weights, value_states)
105
+ attn_output = attn_output.transpose(1, 2).contiguous()
106
+ else:
107
+ attn_output, attn_weights = flash_sigmoid_func(
108
+ query,
109
+ key_states,
110
+ value_states,
111
+ softmax_scale=scaling,
112
+ dropout_p=dropout,
113
+ return_attn_probs=True,
114
+ )
115
+
116
+ return attn_output, attn_weights
117
+
118
+ def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
119
+ """
120
+ This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
121
+ num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
122
+ """
123
+ batch, num_key_value_heads, slen, head_dim = hidden_states.shape
124
+ if n_rep == 1:
125
+ return hidden_states
126
+ hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)
127
+ return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
128
+
129
+ class PicaPreTrainedModel(PreTrainedModel):
130
+ config_class = PicaConfig
131
+ base_model_prefix = "model"
132
+ supports_gradient_checkpointing = True
133
+ _no_split_modules = ["PicaDecoderLayer"]
134
+ _skip_keys_device_placement = ["past_key_values"]
135
+ _supports_flash_attn_2 = False
136
+ _supports_sdpa = False
137
+ _supports_flex_attn = False
138
+ _supports_cache_class = True
139
+ _supports_quantized_cache = True
140
+ _supports_static_cache = True
141
+ _supports_attention_backend = False
142
+
143
+ def _init_weights(self, module):
144
+ std = self.config.initializer_range
145
+ if isinstance(module, nn.Linear):
146
+ module.weight.data.normal_(mean=0.0, std=std)
147
+ if module.bias is not None:
148
+ module.bias.data.zero_()
149
+ elif isinstance(module, nn.Embedding):
150
+ module.weight.data.normal_(mean=0.0, std=std)
151
+ if module.padding_idx is not None:
152
+ module.weight.data[module.padding_idx].zero_()
153
+ elif isinstance(module, LlamaRMSNorm):
154
+ module.weight.data.fill_(1.0)
155
+
156
+ class PicaModel(PicaPreTrainedModel):
157
+ def __init__(self, config: PicaConfig):
158
+ super().__init__(config)
159
+ self.padding_idx = config.pad_token_id
160
+ self.vocab_size = config.vocab_size
161
+
162
+ self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx)
163
+ self.embed_norm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
164
+ self.layers = nn.ModuleList(
165
+ [PicaDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
166
+ )
167
+ self.norm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
168
+ self.gradient_checkpointing = False
169
+
170
+ # Initialize weights and apply final processing
171
+ self.post_init()
172
+
173
+ def get_input_embeddings(self):
174
+ return self.embed_tokens
175
+
176
+ def set_input_embeddings(self, value):
177
+ self.embed_tokens = value
178
+
179
+ def forward(
180
+ self,
181
+ input_ids: Optional[torch.LongTensor] = None,
182
+ attention_mask: Optional[torch.Tensor] = None,
183
+ position_ids: Optional[torch.LongTensor] = None,
184
+ past_key_values: Optional[Cache] = None,
185
+ inputs_embeds: Optional[torch.FloatTensor] = None,
186
+ use_cache: Optional[bool] = None,
187
+ output_attentions: Optional[bool] = None,
188
+ output_hidden_states: Optional[bool] = None,
189
+ cache_position: Optional[torch.LongTensor] = None,
190
+ **kwargs,
191
+ ) -> BaseModelOutputWithPast:
192
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
193
+ output_hidden_states = (
194
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
195
+ )
196
+ use_cache = use_cache if use_cache is not None else self.config.use_cache
197
+
198
+ if (input_ids is None) ^ (inputs_embeds is not None):
199
+ raise ValueError("You must specify exactly one of input_ids or inputs_embeds")
200
+
201
+ if self.gradient_checkpointing and self.training and use_cache:
202
+ logger.warning_once(
203
+ "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`."
204
+ )
205
+ use_cache = False
206
+
207
+ # TODO (joao): remove this exception in v4.56 -- it exists for users that try to pass a legacy cache
208
+ if not isinstance(past_key_values, (type(None), Cache)):
209
+ raise ValueError("The `past_key_values` should be either a `Cache` object or `None`.")
210
+
211
+ if inputs_embeds is None:
212
+ inputs_embeds = self.embed_tokens(input_ids)
213
+
214
+ if use_cache and past_key_values is None:
215
+ past_key_values = DynamicCache()
216
+
217
+ if cache_position is None:
218
+ past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
219
+ cache_position = torch.arange(
220
+ past_seen_tokens, past_seen_tokens + inputs_embeds.shape[1], device=inputs_embeds.device
221
+ )
222
+
223
+ if position_ids is None:
224
+ position_ids = cache_position.unsqueeze(0)
225
+
226
+ causal_mask = self._update_causal_mask(
227
+ attention_mask, inputs_embeds, cache_position, past_key_values, output_attentions
228
+ )
229
+
230
+ hidden_states = self.embed_norm(inputs_embeds)
231
+
232
+ # decoder layers
233
+ all_hidden_states = () if output_hidden_states else None
234
+ all_self_attns = () if output_attentions else None
235
+
236
+ for decoder_layer in self.layers[: self.config.num_hidden_layers]:
237
+ if output_hidden_states:
238
+ all_hidden_states += (hidden_states,)
239
+
240
+ layer_outputs = decoder_layer(
241
+ hidden_states,
242
+ attention_mask=causal_mask,
243
+ position_ids=position_ids,
244
+ past_key_value=past_key_values,
245
+ output_attentions=output_attentions,
246
+ use_cache=use_cache,
247
+ cache_position=cache_position,
248
+ **kwargs,
249
+ )
250
+
251
+ hidden_states = layer_outputs[0]
252
+
253
+ if output_attentions:
254
+ all_self_attns += (layer_outputs[1],)
255
+
256
+ hidden_states = self.norm(hidden_states)
257
+
258
+ # add hidden states from the last decoder layer
259
+ if output_hidden_states:
260
+ all_hidden_states += (hidden_states,)
261
+
262
+ return BaseModelOutputWithPast(
263
+ last_hidden_state=hidden_states,
264
+ past_key_values=past_key_values if use_cache else None,
265
+ hidden_states=all_hidden_states,
266
+ attentions=all_self_attns,
267
+ )
268
+
269
+ def _update_causal_mask(
270
+ self,
271
+ attention_mask: torch.Tensor,
272
+ input_tensor: torch.Tensor,
273
+ cache_position: torch.Tensor,
274
+ past_key_values: Cache,
275
+ output_attentions: bool = False,
276
+ ):
277
+ if flash_sigmoid_func is not None:
278
+ if attention_mask is not None and (attention_mask == 0.0).any():
279
+ return attention_mask
280
+ return None
281
+
282
+ past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
283
+ using_static_cache = isinstance(past_key_values, StaticCache)
284
+ dtype, device = input_tensor.dtype, input_tensor.device
285
+ sequence_length = input_tensor.shape[1]
286
+ if using_static_cache:
287
+ target_length = past_key_values.get_max_cache_shape()
288
+ else:
289
+ target_length = (
290
+ attention_mask.shape[-1]
291
+ if isinstance(attention_mask, torch.Tensor)
292
+ else past_seen_tokens + sequence_length + 1
293
+ )
294
+
295
+ # In case the provided `attention` mask is 2D, we generate a causal mask here (4D).
296
+ causal_mask = LlamaModel._prepare_4d_causal_attention_mask_with_cache_position(
297
+ attention_mask,
298
+ sequence_length=sequence_length,
299
+ target_length=target_length,
300
+ dtype=dtype,
301
+ device=device,
302
+ cache_position=cache_position,
303
+ batch_size=input_tensor.shape[0],
304
+ )
305
+ return causal_mask
306
+
307
+ class PicaDecoderLayer(nn.Module):
308
+ def __init__(self, config: PicaConfig, layer_idx: int):
309
+ super().__init__()
310
+ self.hidden_size = config.hidden_size
311
+
312
+ self.self_attn = PicaAttention(config=config, layer_idx=layer_idx)
313
+
314
+ self.mlp = LlamaMLP(config)
315
+ self.input_layernorm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
316
+ self.post_attention_layernorm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
317
+
318
+ def forward(
319
+ self,
320
+ hidden_states: torch.Tensor,
321
+ attention_mask: Optional[torch.Tensor] = None,
322
+ past_key_value: Optional[Cache] = None,
323
+ output_attentions: Optional[bool] = False,
324
+ use_cache: Optional[bool] = False,
325
+ cache_position: Optional[torch.LongTensor] = None,
326
+ position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # necessary, but kept here for BC
327
+ **kwargs,
328
+ ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:
329
+ residual = hidden_states
330
+
331
+ hidden_states = self.input_layernorm(hidden_states)
332
+
333
+ # Self Attention
334
+ hidden_states, self_attn_weights = self.self_attn(
335
+ hidden_states=hidden_states,
336
+ attention_mask=attention_mask,
337
+ past_key_value=past_key_value,
338
+ output_attentions=output_attentions,
339
+ use_cache=use_cache,
340
+ cache_position=cache_position,
341
+ position_embeddings=position_embeddings,
342
+ **kwargs,
343
+ )
344
+ hidden_states = residual + hidden_states
345
+
346
+ # Fully Connected
347
+ residual = hidden_states
348
+ hidden_states = self.post_attention_layernorm(hidden_states)
349
+ hidden_states = self.mlp(hidden_states)
350
+ hidden_states = residual + hidden_states
351
+
352
+ outputs = (hidden_states,)
353
+ if output_attentions:
354
+ outputs += (self_attn_weights,)
355
+
356
+ return outputs
357
+
358
+
359
+ class KwargsForCausalLM(LossKwargs): ...
360
+
361
+ class PicaForCausalLM(PicaPreTrainedModel, GenerationMixin):
362
+ _tied_weights_keys = ["lm_head.weight"]
363
+ _tp_plan = {"lm_head": "colwise_rep"}
364
+ _pp_plan = {"lm_head": (["hidden_states"], ["logits"])}
365
+
366
+ def __init__(self, config):
367
+ super().__init__(config)
368
+ self.model = PicaModel(config)
369
+ self.vocab_size = config.vocab_size
370
+ self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
371
+
372
+ # Initialize weights and apply final processing
373
+ self.post_init()
374
+
375
+ def get_input_embeddings(self):
376
+ return self.model.embed_tokens
377
+
378
+ def set_input_embeddings(self, value):
379
+ self.model.embed_tokens = value
380
+
381
+ def get_output_embeddings(self):
382
+ return self.lm_head
383
+
384
+ def set_output_embeddings(self, new_embeddings):
385
+ self.lm_head = new_embeddings
386
+
387
+ def set_decoder(self, decoder):
388
+ self.model = decoder
389
+
390
+ def get_decoder(self):
391
+ return self.model
392
+
393
+ def forward(
394
+ self,
395
+ input_ids: Optional[torch.LongTensor] = None,
396
+ attention_mask: Optional[torch.Tensor] = None,
397
+ position_ids: Optional[torch.LongTensor] = None,
398
+ past_key_values: Optional[Cache] = None,
399
+ inputs_embeds: Optional[torch.FloatTensor] = None,
400
+ labels: Optional[torch.LongTensor] = None,
401
+ use_cache: Optional[bool] = None,
402
+ output_attentions: Optional[bool] = None,
403
+ output_hidden_states: Optional[bool] = None,
404
+ cache_position: Optional[torch.LongTensor] = None,
405
+ logits_to_keep: Union[int, torch.Tensor] = 0,
406
+ **kwargs: Unpack[KwargsForCausalLM],
407
+ ) -> CausalLMOutputWithPast:
408
+ r"""
409
+ labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
410
+ Labels for computing the masked language modeling loss. Indices should either be in `[0, ...,
411
+ config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored
412
+ (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
413
+
414
+ logits_to_keep (`int` or `torch.Tensor`, *optional*):
415
+ If an `int`, compute logits for the last `logits_to_keep` tokens. If `0`, calculate logits for all
416
+ `input_ids` (special case). Only last token logits are needed for generation, and calculating them only for that
417
+ token can save memory, which becomes pretty significant for long sequences or large vocabulary size.
418
+ If a `torch.Tensor`, must be 1D corresponding to the indices to keep in the sequence length dimension.
419
+ This is useful when using packed tensor format (single dimension for batch and sequence length).
420
+
421
+ Returns:
422
+
423
+ Example:
424
+
425
+ ```python
426
+ >>> from transformers import AutoTokenizer, LlamaForCausalLM
427
+
428
+ >>> model = LlamaForCausalLM.from_pretrained("meta-llama/Llama-2-7b-hf")
429
+ >>> tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-2-7b-hf")
430
+
431
+ >>> prompt = "Hey, are you conscious? Can you talk to me?"
432
+ >>> inputs = tokenizer(prompt, return_tensors="pt")
433
+
434
+ >>> # Generate
435
+ >>> generate_ids = model.generate(inputs.input_ids, max_length=30)
436
+ >>> tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0]
437
+ "Hey, are you conscious? Can you talk to me?\nI'm not conscious, but I can talk to you."
438
+ ```"""
439
+ output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
440
+ output_hidden_states = (
441
+ output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
442
+ )
443
+
444
+ # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn)
445
+ outputs: BaseModelOutputWithPast = self.model(
446
+ input_ids=input_ids,
447
+ attention_mask=attention_mask,
448
+ position_ids=position_ids,
449
+ past_key_values=past_key_values,
450
+ inputs_embeds=inputs_embeds,
451
+ use_cache=use_cache,
452
+ output_attentions=output_attentions,
453
+ output_hidden_states=output_hidden_states,
454
+ cache_position=cache_position,
455
+ **kwargs,
456
+ )
457
+
458
+ hidden_states = outputs.last_hidden_state
459
+ # Only compute necessary logits, and do not upcast them to float if we are not computing the loss
460
+ slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep
461
+ logits = self.lm_head(hidden_states[:, slice_indices, :])
462
+
463
+ loss = None
464
+ if labels is not None:
465
+ loss = self.loss_function(logits=logits, labels=labels, vocab_size=self.config.vocab_size, **kwargs)
466
+
467
+ return CausalLMOutputWithPast(
468
+ loss=loss,
469
+ logits=logits,
470
+ past_key_values=outputs.past_key_values,
471
+ hidden_states=outputs.hidden_states,
472
+ attentions=outputs.attentions,
473
+ )
474
+
475
+ AutoConfig.register("pica", PicaConfig)
476
+ AutoModel.register(PicaConfig, PicaModel)
477
+ AutoModelForCausalLM.register(PicaConfig, PicaForCausalLM)