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)