VDrontV3-Mini / use.py
MishaGGG's picture
Upload 12 files
769a891 verified
Raw
History Blame Contribute Delete
4.51 kB
# use.py
import os
import json
import torch
import torch.nn.functional as F
from transformers import GPT2TokenizerFast, GPT2Config
from safetensors.torch import load_file
from model import VDrontModel
CONFIG = {
"model_dir": "./VDrontV3-Mini",
"temperature": 0.4,
"top_k": 50,
"max_new_tokens": 200,
"repetition_penalty": 1.2,
"user_token": "<|user|>",
"assistant_token": "<|assistant|>",
}
def format_prompt(user_input):
return f"{CONFIG['user_token']}{user_input}{CONFIG['assistant_token']}"
def main():
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
tokenizer = GPT2TokenizerFast.from_pretrained(CONFIG['model_dir'])
vocab_size = len(tokenizer)
special_tokens = [CONFIG['user_token'], CONFIG['assistant_token']]
tokenizer.add_special_tokens({'additional_special_tokens': special_tokens})
with open(os.path.join(CONFIG['model_dir'], 'architecture.json')) as f:
arch = json.load(f)
config = GPT2Config(
vocab_size=vocab_size,
n_embd=arch['n_embd'],
n_head=arch['n_head'],
n_layer=arch['n_layer'],
n_positions=arch['n_positions'],
layer_norm_epsilon=1e-5,
)
model = VDrontModel(
config=config,
expert_start=arch['expert_start'],
expert_end=arch['expert_end'],
output_index=arch['output_index'],
num_experts=arch['num_experts'],
num_output_versions=arch['num_output_versions'],
)
state = load_file(os.path.join(CONFIG['model_dir'], 'model.safetensors'))
model.load_state_dict(state)
model.to(device)
model.eval()
if model.embed_tokens.num_embeddings < len(tokenizer):
old_embed = model.embed_tokens
new_embed = torch.nn.Embedding(len(tokenizer), old_embed.embedding_dim).to(device)
new_embed.weight.data[:old_embed.num_embeddings] = old_embed.weight.data.to(device)
model.embed_tokens = new_embed
old_lm_head = model.lm_head
new_lm_head = torch.nn.Linear(old_lm_head.in_features, len(tokenizer), bias=False).to(device)
new_lm_head.weight.data[:old_lm_head.out_features] = old_lm_head.weight.data.to(device)
model.lm_head = new_lm_head
model.config.vocab_size = len(tokenizer)
while True:
try:
output_ver = int(input("Mode (0 - base (bad, little answer), 1 - qualitative (normal, medium answer): "))
if output_ver in [0, 1]:
model.set_output_version(output_ver)
break
except ValueError:
pass
print("Chat is ready. Type 'exit' to quit.")
while True:
user_input = input("You: ")
if user_input.lower() in ['exit', 'quit']:
break
prompt = format_prompt(user_input)
input_ids = tokenizer.encode(prompt, return_tensors='pt').to(device)
generated_tokens = []
eos_id = tokenizer.eos_token_id
with torch.no_grad():
for _ in range(CONFIG['max_new_tokens']):
pos = torch.arange(0, input_ids.size(1), device=device).unsqueeze(0)
x = model.embed_tokens(input_ids) + model.embed_positions(pos)
router_logits = model.router(x.mean(dim=1))
expert_idx = router_logits.argmax(dim=-1).item()
model.set_expert_version(expert_idx)
idx_cond = input_ids[:, -model.config.n_positions:]
logits, _ = model(idx_cond)
logits = logits[:, -1, :] / CONFIG['temperature']
for token_id in set(input_ids[0].tolist()):
logits[0, token_id] /= CONFIG['repetition_penalty']
if CONFIG['top_k'] is not None:
v, _ = torch.topk(logits, min(CONFIG['top_k'], logits.size(-1)))
logits[logits < v[:, [-1]]] = -float('Inf')
probs = F.softmax(logits, dim=-1)
idx_next = torch.multinomial(probs, num_samples=1)
next_token = idx_next.item()
if next_token == eos_id:
break
generated_tokens.append(next_token)
input_ids = torch.cat((input_ids, idx_next), dim=1)
full_text = tokenizer.decode(generated_tokens, skip_special_tokens=True)
print(f"AI: {full_text}")
if __name__ == '__main__':
main()