from __future__ import annotations from typing import Optional import torch from torch import nn from transformers import PreTrainedModel from transformers.modeling_outputs import CausalLMOutputWithPast from .configuration_aurora import AuroraHFConfig from .aurora_config import AuroraConfig as NativeAuroraConfig from .aurora_model import AuroraForCausalLM as NativeAuroraForCausalLM class AuroraForCausalLM(PreTrainedModel): config_class = AuroraHFConfig base_model_prefix = "aurora" main_input_name = "input_ids" _supports_cache_class = False _tied_weights_keys = {} all_tied_weights_keys = {} def __init__(self, config: AuroraHFConfig): super().__init__(config) native_config = NativeAuroraConfig( model_name=config.model_name, vocab_size=config.vocab_size, hidden_size=config.hidden_size, num_layers=config.num_layers, num_attention_heads=config.num_attention_heads, num_key_value_heads=config.num_key_value_heads, intermediate_size=config.intermediate_size, context_length=config.context_length, rope_theta=config.rope_theta, rms_norm_eps=config.rms_norm_eps, qk_norm=config.qk_norm, tie_word_embeddings=config.tie_word_embeddings, attention_bias=config.attention_bias, mlp_bias=config.mlp_bias, dropout=config.dropout, num_experts=config.num_experts, router_aux_loss_coef=config.router_aux_loss_coef, router_z_loss_coef=config.router_z_loss_coef, router_noise_scale=config.router_noise_scale, moe_capacity_factor=config.moe_capacity_factor, router_use_gate_weight=config.router_use_gate_weight, ) self.aurora = NativeAuroraForCausalLM(native_config) def load_state_dict(self, state_dict, strict: bool = True, assign: bool = False, **kwargs): # The raw Aurora export stores the native module at the repository root; # the Transformers wrapper stores it below `aurora`. remapped = { (key if key.startswith("aurora.") else f"aurora.{key}"): value for key, value in state_dict.items() } return super().load_state_dict(remapped, strict=strict, assign=assign, **kwargs) def get_input_embeddings(self): return self.aurora.embed_tokens def set_input_embeddings(self, value): self.aurora.embed_tokens = value def get_output_embeddings(self): return self.aurora.lm_head def set_output_embeddings(self, new_embeddings): self.aurora.lm_head = new_embeddings def prepare_inputs_for_generation(self, input_ids, **kwargs): return {"input_ids": input_ids} def forward( self, input_ids: torch.LongTensor, attention_mask: Optional[torch.Tensor] = None, labels: Optional[torch.LongTensor] = None, **kwargs, ) -> CausalLMOutputWithPast: logits, loss = self.aurora(input_ids=input_ids, labels=labels) return CausalLMOutputWithPast(loss=loss, logits=logits)