|
|
| 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() |