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)