Views
No views yet

1pip install -r requirements.txt
2pip install transformers datasets accelerate safetensors1import torch
2from transformers import AutoTokenizer, AutoModelForCausalLM
3
4
5device = "cuda" if torch.cuda.is_available() else "cpu"
6
7model_path = "DireDreadlord/GemCod-Zircon-270M-XM"
8
9tokenizer = AutoTokenizer.from_pretrained(model_path)
10
11model = AutoModelForCausalLM.from_pretrained(model_path)
12model.to(device)
13model.eval()
14
15if tokenizer.pad_token is None:
16 tokenizer.pad_token = tokenizer.eos_token
17
18model.resize_token_embeddings(len(tokenizer))
19
20
21chat_template = """{% for message in messages %}{% if message['role'] == 'user' %}User: {{ message['content'] }}
22{% elif message['role'] == 'assistant' %}Assistant: {{ message['content'] }}
23{% endif %}{% endfor %}"""
24tokenizer.chat_template = chat_template
25
26def generate_code(prompt, max_tokens=512) -> str:
27 messages = [
28 {
29 "role": "user",
30 "content": prompt
31 }
32 ]
33
34 formatted_prompt = tokenizer.apply_chat_template(
35 messages,
36 tokenize=False,
37 add_generation_prompt=True
38 )
39
40 inputs = tokenizer(formatted_prompt, return_tensors="pt").to(device)
41 input_length = inputs["input_ids"].shape[1]
42
43 with torch.no_grad():
44 outputs = model.generate(
45 **inputs,
46 max_new_tokens=max_tokens,
47 do_sample=False,
48 num_beams=1,
49 pad_token_id=tokenizer.eos_token_id,
50 eos_token_id=tokenizer.eos_token_id,
51 use_cache=False,
52 )
53
54 generated_tokens = outputs[0][input_length:]
55 generated_text = tokenizer.decode(generated_tokens, skip_special_tokens=True)
56 return generated_text
57
58
59def test_fim(prefix, suffix, max_tokens=32):
60 user = (
61 "Complete the missing middle between the markers.\n"
62 "PREFIX:\n" + prefix + "\n\n<FILL>\n\nSUFFIX:\n" + suffix
63 )
64 return generate_code(user, max_tokens=max_tokens)
65
66
67def test_multiturn_chat(history, question, max_tokens=128):
68 messages = []
69 for role, content in history:
70 messages.append({"role": role, "content": content})
71 messages.append({"role": "user", "content": question})
72
73 formatted = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
74 inputs = tokenizer(formatted, return_tensors="pt").to(device)
75 input_length = inputs["input_ids"].shape[1]
76
77 with torch.no_grad():
78 outputs = model.generate(
79 **inputs,
80 max_new_tokens=max_tokens,
81 do_sample=False,
82 num_beams=1,
83 pad_token_id=tokenizer.eos_token_id,
84 eos_token_id=tokenizer.eos_token_id,
85 use_cache=False,
86 )
87
88 gen = outputs[0][input_length:]
89 return tokenizer.decode(gen, skip_special_tokens=True)
90
91
92if __name__ == "__main__":
93 # FIM example
94 pref = "def factorial(n):"
95 suf = "return n"
96 print("FIM test:\n", test_fim(pref, suf))
97
98 # multi-turn chat example
99 history = [("user", "Help me write a function to implement an bubble sort algorithm."), ("assistant", "Sure - what language?")]
100 print("Chat test:\n", test_multiturn_chat(history, "write it in python"))