Sol-Lite-2 / modeling_sol2.py
j0no12's picture
Upload final Sol Lite 2 with Sol2ForCausalLM and banner
53d1439 verified
Raw History Blame Contribute Delete
2.68 kB
import torch
from torch.nn import functional as F
from transformers import PreTrainedModel, GenerationMixin
from transformers.modeling_outputs import CausalLMOutput
from .configuration_sol2 import Sol2Config
from .sol2_core import Config, SolLite2
class Sol2ForCausalLM(PreTrainedModel, GenerationMixin):
config_class=Sol2Config
base_model_prefix="model"
main_input_name="input_ids"
_no_split_modules=["Block"]
def __init__(self,config):
super().__init__(config)
keys=("vocab_size","width","heads","kv_heads","blocks","ffn_width",
"recurrent_start","recurrent_blocks","passes","loop_conditioning",
"rope_theta","max_context","backend")
self.model=SolLite2(Config(**{k:getattr(config,k) for k in keys}))
self.post_init()
def get_input_embeddings(self): return self.model.embedding
def set_input_embeddings(self,value): self.model.embedding=value
def get_output_embeddings(self): return self.model.embedding
def set_output_embeddings(self,value): self.model.embedding=value
def prepare_inputs_for_generation(self,input_ids,attention_mask=None,**kwargs):
return {"input_ids":input_ids,"attention_mask":attention_mask,"use_cache":False}
def forward(self,input_ids=None,attention_mask=None,labels=None,
return_dict=None,use_cache=False,past_key_values=None,**kwargs):
if input_ids is None: raise ValueError("input_ids are required")
if past_key_values is not None: raise ValueError("KV caching is not implemented")
if attention_mask is None or bool(attention_mask.bool().all()):
logits=self.model(input_ids)
else:
# Removing padding preserves RoPE positions for each sequence.
rows=[]
for ids,mask in zip(input_ids,attention_mask):
keep=mask.bool();positions=keep.nonzero(as_tuple=True)[0]
if positions.numel()==0: raise ValueError("empty attention mask")
values=self.model(ids[keep][None])[0]
row=values.new_zeros((ids.numel(),self.config.vocab_size))
rows.append(row.index_copy(0,positions,values))
logits=torch.stack(rows)
loss=None
if labels is not None:
targets=labels[:,1:].clone()
if attention_mask is not None: targets.masked_fill_(~attention_mask[:,1:].bool(),-100)
loss=F.cross_entropy(logits[:,:-1].float().reshape(-1,logits.shape[-1]),targets.reshape(-1),ignore_index=-100)
if return_dict is False: return ((loss,) if loss is not None else ())+ (logits,)
return CausalLMOutput(loss=loss,logits=logits)