| 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): |
| |
| |
| 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) |
|
|