Atomight-V2.2-Expt-Hybrid-1.8B / modeling_hybrid.py
NovatasticRoScript's picture
Upload folder using huggingface_hub
f4927a9 verified
Raw
History Blame Contribute Delete
3.21 kB
import torch
import torch.nn as nn
from transformers import PreTrainedModel, AutoModelForCausalLM, AutoConfig, GenerationMixin
from transformers.modeling_outputs import CausalLMOutputWithPast
from configuration_hybrid import HybridConfig
class HybridForCausalLM(PreTrainedModel, GenerationMixin):
config_class = HybridConfig
base_model_prefix = "hybrid"
def __init__(self, config, load_weights=False):
super().__init__(config)
qwen_obj_config = AutoConfig.from_pretrained(config.qwen_model_name)
minicpm_obj_config = AutoConfig.from_pretrained(config.minicpm_model_name)
if load_weights:
self.qwen = AutoModelForCausalLM.from_pretrained(config.qwen_model_name, torch_dtype=torch.bfloat16, low_cpu_mem_usage=True)
self.minicpm = AutoModelForCausalLM.from_pretrained(config.minicpm_model_name, torch_dtype=torch.bfloat16, low_cpu_mem_usage=True)
else:
self.qwen = AutoModelForCausalLM.from_config(qwen_obj_config)
self.minicpm = AutoModelForCausalLM.from_config(minicpm_obj_config)
self.qwen.eval()
self.minicpm.eval()
for p in self.qwen.parameters(): p.requires_grad = False
for p in self.minicpm.parameters(): p.requires_grad = False
self.embed_projection = nn.Linear(qwen_obj_config.hidden_size, minicpm_obj_config.hidden_size, bias=False, dtype=torch.bfloat16)
self.minicpm_to_qwen = nn.Linear(minicpm_obj_config.hidden_size, qwen_obj_config.hidden_size, bias=False, dtype=torch.bfloat16)
self.mix_gate = nn.Parameter(torch.tensor(0.0, dtype=torch.bfloat16))
self.config.use_cache = False
def forward(self, input_ids=None, attention_mask=None, labels=None, past_key_values=None, use_cache=None, **kwargs):
with torch.no_grad():
qwen_outputs = self.qwen(input_ids=input_ids, attention_mask=attention_mask, use_cache=False, output_hidden_states=True, return_dict=True, **kwargs)
qwen_hidden = qwen_outputs.hidden_states[-1]
qwen_embeds = self.qwen.get_input_embeddings()(input_ids)
minicpm_embeds = self.embed_projection(qwen_embeds)
minicpm_outputs = self.minicpm(inputs_embeds=minicpm_embeds, attention_mask=attention_mask, use_cache=False, output_hidden_states=True, return_dict=True, **kwargs)
minicpm_hidden = minicpm_outputs.hidden_states[-1]
fused_hidden = qwen_hidden + torch.sigmoid(self.mix_gate) * self.minicpm_to_qwen(minicpm_hidden)
lm_head = self.qwen.get_output_embeddings() or self.qwen.lm_head
final_logits = lm_head(fused_hidden)
loss = None
if labels is not None:
loss_fct = nn.CrossEntropyLoss(ignore_index=-100)
loss = loss_fct(final_logits[:, :-1, :].contiguous().reshape(-1, final_logits.size(-1)), labels[:, 1:].contiguous().reshape(-1))
return CausalLMOutputWithPast(loss=loss, logits=final_logits)
def prepare_inputs_for_generation(self, input_ids, past_key_values=None, attention_mask=None, **kwargs):
return {"input_ids": input_ids[:, -1:] if past_key_values is not None else input_ids, "attention_mask": attention_mask, "past_key_values": None}