VISTA-24M / modeling_dense.py
AwakeningOS's picture
Release VISTA-24M: model, architecture diagrams, training recipe and evaluation evidence
9287d39 verified
Raw History Blame Contribute Delete
3.26 kB
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)