"""Hugging Face model wrapper for TaoNet.""" from torch import nn from transformers import GenerationMixin, PreTrainedModel from transformers.modeling_outputs import CausalLMOutput try: from .configuration_taonet import TaoNetConfig from .taonet_model import SimpleLLM, build_runtime_config except ImportError: from configuration_taonet import TaoNetConfig from taonet_model import SimpleLLM, build_runtime_config class TaoNetForCausalLM(PreTrainedModel, GenerationMixin): """Transformers-compatible TaoNet causal LM.""" config_class = TaoNetConfig base_model_prefix = "model" supports_gradient_checkpointing = False def __init__(self, config): super().__init__(config) runtime_config = build_runtime_config(config) self.model = SimpleLLM(runtime_config) self.post_init() self.tie_weights() def get_input_embeddings(self): if getattr(self.model, "use_factorized_embedding", False): return self.model.token_embedding.embed return self.model.token_embedding def set_input_embeddings(self, value): if getattr(self.model, "use_factorized_embedding", False): self.model.token_embedding.embed = value else: self.model.token_embedding = value def get_output_embeddings(self): return self.model.output_head def set_output_embeddings(self, new_embeddings): self.model.output_head = new_embeddings def tie_weights(self, *args, **kwargs): del args, kwargs if not getattr(self.model, "use_factorized_embedding", False): self.model.output_head.weight = self.get_input_embeddings().weight def _init_weights(self, module): if isinstance(module, nn.Linear): nn.init.normal_(module.weight, mean=0.0, std=self.config.init_std) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=self.config.init_std) def forward( self, input_ids=None, attention_mask=None, labels=None, inputs_embeds=None, return_dict=None, **kwargs, ): del kwargs return_dict = return_dict if return_dict is not None else self.config.use_return_dict outputs = self.model( input_ids=input_ids, attention_mask=attention_mask, labels=None, inputs_embeds=inputs_embeds, ) loss = None if labels is not None: shift_logits = outputs["logits"][..., :-1, :].contiguous() shift_labels = labels[..., 1:].contiguous() loss_fct = nn.CrossEntropyLoss(ignore_index=-100) loss = loss_fct( shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1), ) if not return_dict: return (loss, outputs["logits"]) return CausalLMOutput(loss=loss, logits=outputs["logits"]) def prepare_inputs_for_generation(self, input_ids, attention_mask=None, **kwargs): return {"input_ids": input_ids, "attention_mask": attention_mask}