File size: 3,261 Bytes
9287d39
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
from contextlib import nullcontext
import torch
from torch.nn import functional as F
from transformers import PreTrainedModel
from transformers.generation import GenerationMixin
from transformers.modeling_outputs import BaseModelOutput,CausalLMOutput
from .configuration_dense import ModernDenseConfig
from .dense import Config,DenseLM,prepare_layout

class ModernDensePreTrainedModel(PreTrainedModel):
    config_class=ModernDenseConfig
    base_model_prefix='model'
    supports_gradient_checkpointing=False
    def _init_weights(self,module):return

class ModernDenseModel(ModernDensePreTrainedModel):
    def __init__(self,config):
        super().__init__(config);self.dense=DenseLM(Config(**config.dense_config));self.post_init()
    def get_input_embeddings(self):return self.dense.embed
    def forward(self,input_ids,attention_mask=None,output_hidden_states=None,return_dict=None,**kwargs):
        if attention_mask is None:attention_mask=torch.ones_like(input_ids)
        layout=prepare_layout(attention_mask.detach().to('cpu',dtype=torch.int32),input_ids.device,self.dense.cfg.backend)
        context=torch.autocast('cuda',dtype=torch.bfloat16) if input_ids.is_cuda else nullcontext()
        with context:hidden=self.dense.hidden(input_ids,layout)
        states=(hidden,) if output_hidden_states else None
        if return_dict is False:return (hidden,states) if states else (hidden,)
        return BaseModelOutput(last_hidden_state=hidden,hidden_states=states)

class ModernDenseForCausalLM(ModernDensePreTrainedModel,GenerationMixin):
    def __init__(self,config):
        super().__init__(config);self.model=ModernDenseModel(config);self.post_init()
    def get_input_embeddings(self):return self.model.dense.embed
    def set_input_embeddings(self,value):self.model.dense.embed=value
    def get_output_embeddings(self):return self.model.dense.lm_head
    def set_output_embeddings(self,value):self.model.dense.lm_head=value
    def prepare_inputs_for_generation(self,input_ids,attention_mask=None,**kwargs):return {'input_ids':input_ids,'attention_mask':attention_mask}
    def forward(self,input_ids,attention_mask=None,labels=None,output_hidden_states=None,return_dict=None,**kwargs):
        output=self.model(input_ids,attention_mask,output_hidden_states,True)
        context=torch.autocast('cuda',dtype=torch.bfloat16) if input_ids.is_cuda else nullcontext()
        with context:logits=self.model.dense.lm_head(output.last_hidden_state)
        loss=None
        if labels is not None:
            shifted=labels[:,1:].contiguous();pred=logits[:,:-1].contiguous()
            # A target is valid only when its predictor and target positions belong
            # to the same non-padding segment. This also fixes HF left padding.
            if attention_mask is not None:
                left=attention_mask[:,:-1];right=attention_mask[:,1:]
                valid=(left>0)&(left==right)
                shifted=shifted.masked_fill(~valid,-100)
            loss=F.cross_entropy(pred.float().view(-1,pred.shape[-1]),shifted.view(-1),ignore_index=-100)
        if return_dict is False:return ((loss,logits) if loss is not None else (logits,))
        return CausalLMOutput(loss=loss,logits=logits,hidden_states=output.hidden_states)