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}