import torch, torch.nn as nn
from transformers import AutoModelForCausalLM, AutoTokenizer
from transformers.models.llama.modeling_llama import LlamaDecoderLayer, LlamaRMSNorm
from huggingface_hub import hf_hub_download
REPO = "ewinregirgojr/minicpm5-stock-analyst-v2-mtp"
backbone = AutoModelForCausalLM.from_pretrained(REPO, torch_dtype=torch.bfloat16).eval()
tokenizer = AutoTokenizer.from_pretrained(REPO)
config = backbone.config
class MTPHead(nn.Module):
def __init__(self, config):
super().__init__()
self.input_proj = nn.Linear(2 * config.hidden_size, config.hidden_size, bias=False)
self.pre_norm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.decoder_layer = LlamaDecoderLayer(config, layer_idx=config.num_hidden_layers)
self.out_norm = LlamaRMSNorm(config.hidden_size, eps=config.rms_norm_eps)
def forward(self, prev_hidden, token_embed, position_ids, position_embeddings):
x = self.pre_norm(self.input_proj(torch.cat([prev_hidden, token_embed], dim=-1)))
out = self.decoder_layer(x.unsqueeze(1), position_ids=position_ids.unsqueeze(1),
position_embeddings=position_embeddings)
if isinstance(out, tuple):
out = out[0]
return self.out_norm(out.squeeze(1))
head = MTPHead(config).to(torch.bfloat16)
head.load_state_dict(torch.load(hf_hub_download(REPO, "mtp_head.pt"), map_location="cpu"), strict=True)
head.eval()