| 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} |
|
|