Aurora-Proelia-ChatML / modeling_aurora.py
arthu1's picture
Fix normal Transformers inference loading
01cfd4f verified
Raw
History Blame Contribute Delete
3.16 kB
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)