xgmx / modeling.py
linzhe's picture
Upload modeling.py with huggingface_hub
bf73969 verified
Raw
History Blame Contribute Delete
2.93 kB
"""人格预测 v2 - 单模型多头"""
import json, numpy as np, torch, torch.nn as nn
EMBEDDING_DIM, VEC_DIM = 768, 64
TRUNK = [256, 128, 64]
HEADS = {
"ocean": (["开放性","尽责性","外向性","宜人性","神经质"], 10),
"four": (["力量型","活泼型","完美型","和平型"], 100),
"color": (["红","蓝","黄","绿"], 100),
"mbti": (["E","I","S","N","T","F","J","P"], 100),
"enneagram": ([f"type_{i}" for i in range(1,10)], 100),
}
class PersonalityModel(nn.Module):
def __init__(self):
super().__init__()
layers=[]; prev=EMBEDDING_DIM
for h in TRUNK: layers+=[nn.Linear(prev,h),nn.BatchNorm1d(h),nn.ReLU(),nn.Dropout(0.2)]; prev=h
self.trunk=nn.Sequential(*layers)
self.heads=nn.ModuleDict({n:nn.Sequential(nn.Linear(VEC_DIM,len(d)),nn.Sigmoid()) for n,(d,_) in HEADS.items()})
def forward(self,x):
v=self.trunk(x); return {n:h(v) for n,h in self.heads.items()}, v
class PersonalityPredictor:
def __init__(self, model_dir=".", device=None):
self.device=device or ("cuda"if torch.cuda.is_available()else"cpu")
from sentence_transformers import SentenceTransformer
self.enc=SentenceTransformer("shibing624/text2vec-base-chinese",device=self.device)
self.model=PersonalityModel().to(self.device)
ckpt=torch.load(f"{model_dir}/personality_model.pt",map_location=self.device)
self.model.load_state_dict(ckpt["state_dict"]); self.model.eval()
def predict(self, text):
emb=self.enc.encode([text],normalize_embeddings=True)
x=torch.tensor(emb,dtype=torch.float32).to(self.device)
with torch.no_grad(): out,vec=self.model(x)
result={}
for n,(dims,scale) in HEADS.items():
s=out[n].cpu().numpy()[0]; r={}
for i,d in enumerate(dims): r[d]=round(float(s[i])*scale,1)
if n in("four","color"):
r["_primary"]=dims[int(np.argmax(s))]
elif n=="mbti":
e,i_,s_,n_,t,f,j,p=s
r["_mbti"]=("E"if e>i_ else"I")+("S"if s_>n_ else"N")+("T"if t>f else"F")+("J"if j>p else"P")
elif n=="enneagram":
idx=int(np.argmax(s)); names=["完美型","助人型","成就型","自我型","理智型","疑惑型","活跃型","领袖型","和平型"]
r["_primary"]=f"{idx+1}号·{names[idx]}"
result[n]=r
result["vector"]=vec.cpu().numpy()[0].tolist()
return result
if __name__=="__main__":
p=PersonalityPredictor()
for t in ["曹操性格奸诈多疑、雄才大略,善于用人却心狠手辣",
"此人极其暴躁冲动,一言不合就动手打人,毫无耐心"]:
r=p.predict(t); print(f"\n{t[:30]}...")
print(f" OCEAN:{r['ocean']} 四型:{r['four'].get('_primary')} MBTI:{r['mbti'].get('_mbti')}")