Views
No views yet
1from tokenizers import Tokenizer
2import onnxruntime as ort
3import numpy as np
4
5def generate_chat_response(prompt_input, stopping_text=None, tokenizer_path="tokenizer.json", model_path="model.onnx", *, temperature=0.1, top_k=50, top_p=0.5, repetition_penalty=1.25, max_tokens=4096):
6 """
7 Args:
8 prompt_input (list or str):
9 - If a list, it should be a list of dictionaries representing the chat history.
10 Each dictionary should have "role" and "content" keys.
11 - If a string, it will be treated as a raw user input and converted internally.
12 tokenizer_path (str): Path to the tokenizer JSON file.
13 model_path (str): Path to the ONNX model file.
14
15 Returns:
16 str: The generated response from the model.
17 """
18 # Load the tokenizer and ONNX model
19 tokenizer = Tokenizer.from_file(tokenizer_path)
20 session = ort.InferenceSession(model_path)
21
22 # Handle raw string input by converting it to the expected format
23 if isinstance(prompt_input, str):
24 messages = [{"role": "user", "content": prompt_input}]
25 elif isinstance(prompt_input, list):
26 messages = prompt_input
27 else:
28 raise ValueError("prompt_input must be either a string or a list of dictionaries.")
29
30 # Construct the input text based on the chat schema
31 input_text_parts = []
32 for message in messages:
33 role = message["role"]
34 content = message["content"]
35 if role == "user":
36 input_text_parts.append(f"<instruction>{content}</instruction>")
37 elif role == "assistant":
38 input_text_parts.append(content)
39 input_text = "".join(input_text_parts)
40
41 # Tokenize the input text
42 encoded_input = tokenizer.encode(input_text)
43 input_ids = np.array([encoded_input.ids], dtype=np.int64) # Shape: (1, sequence_length)
44
45 # Prepare attention mask and position IDs
46 attention_mask = np.ones_like(input_ids, dtype=np.int64) # Shape: (1, sequence_length)
47 position_ids = np.arange(input_ids.shape[1], dtype=np.int64).reshape(1, -1) # Shape: (1, sequence_length)
48
49 # Prepare past_key_values (initialize as empty tensors)
50 num_layers = 12 # Adjust based on your model's number of layers
51 num_heads = 12 # Adjust based on your model's number of attention heads
52 head_dim = 64 # Adjust based on your model's head dimensionality
53 batch_size = 1
54 past_sequence_length = 0 # Initially, no past context exists
55
56 past_key_values = []
57 for _ in range(num_layers):
58 key = np.zeros((batch_size, num_heads, past_sequence_length, head_dim), dtype=np.float32)
59 value = np.zeros((batch_size, num_heads, past_sequence_length, head_dim), dtype=np.float32)
60 past_key_values.append(key)
61 past_key_values.append(value)
62
63 eos_token_id = 2
64
65 def sample_next_token(logits, generated_ids, temperature=1.0, top_k=0, top_p=1.0, repetition_penalty=1.0):
66 logits = logits.copy()
67 for token_id in set(generated_ids):
68 if logits[token_id] < 0:
69 logits[token_id] *= repetition_penalty
70 else:
71 logits[token_id] /= repetition_penalty
72
73 logits = logits / temperature
74 exp_logits = np.exp(logits - np.max(logits))
75 probs = exp_logits / exp_logits.sum()
76
77 if top_k > 0:
78 indices_to_remove = probs < np.sort(probs)[-top_k]
79 probs[indices_to_remove] = 0
80 probs = probs / probs.sum()
81
82 if top_p < 1.0:
83 sorted_indices = np.argsort(probs)[::-1]
84 sorted_probs = probs[sorted_indices]
85 cumulative_probs = np.cumsum(sorted_probs)
86 cutoff_index = np.searchsorted(cumulative_probs, top_p)
87 probs[sorted_indices[cutoff_index+1:]] = 0
88 probs = probs / probs.sum()
89
90 next_token_id = np.random.choice(len(probs), p=probs)
91 return next_token_id
92
93 # Autoregressive generation loop
94 generated_token_ids = []
95 for _ in range(max_tokens):
96 input_feed = {
97 "input_ids": input_ids,
98 "attention_mask": attention_mask,
99 "position_ids": position_ids,
100 }
101
102 for i in range(num_layers):
103 input_feed[f"past_key_values.{i}.key"] = past_key_values[2 * i]
104 input_feed[f"past_key_values.{i}.value"] = past_key_values[2 * i + 1]
105
106 outputs = session.run(None, input_feed)
107 logits = outputs[0] # Shape: (1, sequence_length, vocab_size)
108 updated_past_key_values = outputs[1:]
109
110 last_logits = logits[0, -1, :]
111 next_token_id = sample_next_token(
112 last_logits,
113 generated_token_ids,
114 temperature=temperature,
115 top_k=top_k,
116 top_p=top_p,
117 repetition_penalty=repetition_penalty
118 )
119
120 generated_token_ids.append(next_token_id)
121
122 if next_token_id == eos_token_id:
123 break
124
125 if stopping_text:
126 decoded_text = tokenizer.decode(generated_token_ids)
127 if stopping_text in decoded_text:
128 break
129
130 input_ids = np.array([[next_token_id]], dtype=np.int64)
131 attention_mask = np.ones_like(input_ids, dtype=np.int64)
132 position_ids = np.array([[len(encoded_input.ids) + len(generated_token_ids)]], dtype=np.int64)
133 past_key_values = updated_past_key_values
134
135 # Decode the generated token IDs into text
136 generated_text = tokenizer.decode(generated_token_ids)
137
138 # Remove the stop sequence from the generated text
139 if stopping_text:
140 generated_text = generated_text.split(stopping_text)[0]
141
142 return generated_text
143
144# Example usage
145if __name__ == "__main__":
146 # Using raw string input
147 raw_input = "Era uma vez"
148 response = generate_chat_response(raw_input, max_tokens=50)
149 response
150 # Era uma vez um jovem chamado Jack que vivia em uma pequena cidade chamada Little Rock. Ele era conhecido por ser muito inteligente e tinha um senso de humor peculiar.
151 # Jack sempre foi fascinado pela cultura pop, mas nunca imaginou que seria tão interessante para
152
153 # Using list of dictionaries input
154 dict_input = [
155 {"role": "user", "content": "Quais são as melhores práticas para programação em Python?"}
156 ]
157 response = generate_chat_response(dict_input)
158 response
159 # 1. Use uma variedade de bibliotecas e estruturas: Python é um sistema amplamente utilizado para programação em vários idiomas. Ele permite que você crie aplicativos complexos com facilidade, sem a necessidade de escrever código complexo.
160
161 # 2. Minimize o tamanho da pilha: O uso excessivo do espaço na memória pode levar ao desempenho lento. Portanto, minimize o tamanho da sua pilha usando funções como min-stack ou splunk.
162
163 # 3. Implementar tratamento de erros: A implementação adequada dos testes unitários garante que seu aplicativo funcione conforme esperado.
164
165 # 4. Teste seus aplicativos minuciosamente: teste exaustivamente os recursos do seu aplicativo, incluindo as funcionalidades desejadas, antes de implementá-los.
166
167 # 5. Mantenha-se atualizado sobre novas tecnologias: As atualizações mais recentes podem melhorar significativamente suas capacidades de desenvolvimento.
168
169 # 6. Monitore regularmente seu código: monitore continuamente seu código, verificando se há bugs e fazendo alterações.
170
171 # 7. Utilize ferramentas de gerenciamento de projetos: Ferramentas de gerenciamento de projetos, como Trello, Asana e Jira, permitem gerenciar tarefas e prazos de maneira eficiente.
172
173 # 8. Colabore com outros desenvolvedores: Trabalhe com outras equipes e parceiros para colaborar no projeto.
174
175 # 9. Estabelecer relacionamentos sólidos entre diferentes linguagens de programação: Construir relacionamentos fortes entre diferentes linguagens de programação, garantindo compatibilidade e flexibilidade.
176
177 # 10. Considere usar frameworks de terceiros: Frameworks, como Python, Ruby on Rails e Node.js, fornecem APIs robustas para construir aplicações web complexas.
178
179 # Concluindo, a programação em Python requer atenção aos detalhes, mas também envolve muito trabalho duro e dedicação. Seguindo essas práticas recomendadas e utilizando técnicas adequadas de desenvolvimento, você poderá criar aplicativos poderosos e eficientes que atendam às necessidades específicas das empresas.