Yonetix-V1 / use_model.py
taojian4110's picture
Upload use_model.py with huggingface_hub
eb3c72b verified
Raw
History Blame Contribute Delete
1.57 kB
"""Yonetix-V1 快速使用脚本"""
import json
import torch
from model import YonetixTransformer, ModelConfig
# 加载模型
with open("config.json") as f:
config = ModelConfig(**json.load(f))
model = YonetixTransformer(config)
model.load_state_dict(torch.load("model.pt", map_location="cpu", weights_only=True))
model.eval()
print(f"✅ 模型加载完成,参数量: {sum(p.numel() for p in model.parameters()) / 1e6:.2f}M")
# 加载词表
vocab = {}
with open("vocab.txt") as f:
for line in f:
idx, token = line.strip().split("\t")
vocab[int(idx)] = token
PAD, UNK, BOS, EOS = 0, 1, 2, 3
def encode(text, max_len=128):
ids = [BOS]
for ch in text:
found = False
for idx, token in vocab.items():
if idx < 4:
continue
if token == ch:
ids.append(idx)
found = True
break
if not found:
ids.append(UNK)
ids.append(EOS)
return torch.tensor([ids[:max_len]])
def decode(ids):
return "".join(vocab.get(int(i), "") for i in ids if int(i) > 3)
def chat(prompt, max_new_tokens=128):
input_ids = encode(prompt)
output_ids = model.generate(input_ids, max_new_tokens=max_new_tokens)
reply = output_ids[0][input_ids.size(1):].tolist()
return decode(reply)
# 测试
print("\n🤖 Yonetix-V1 测试对话")
print("=" * 40)
tests = [
"你好,你是谁?",
"介绍一下YONETIX",
"你会做什么?",
]
for t in tests:
print(f"\n👤 > {t}")
print(f"🤖 > {chat(t)}")