SubZtep commited on
Commit
47b302d
·
1 Parent(s): 9e8bc15

feat: flexible modal message handling

Browse files
Files changed (4) hide show
  1. .gitignore +1 -0
  2. ai_engine.py +160 -39
  3. api.py +4 -4
  4. requirements.txt +1 -0
.gitignore CHANGED
@@ -1 +1,2 @@
1
  __pycache__
 
 
1
  __pycache__
2
+ venv
ai_engine.py CHANGED
@@ -1,44 +1,165 @@
1
- from transformers import pipeline
 
 
2
 
3
- print("Loading model...")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4
 
5
- generator = pipeline(
6
- "text-generation",
7
- model="Qwen/Qwen2.5-0.5B-Instruct"
8
- # model="google/flan-t5-small"
9
- )
10
 
11
- SYSTEM_PROMPT = (
12
- "You are a helpful assistant. "
13
- "Answer clearly and concisely."
14
- )
15
 
16
- def build_prompt(history: list[str], message: str) -> str:
17
- prompt = SYSTEM_PROMPT + "\n\n"
 
18
 
19
- for line in history:
20
- prompt += line + "\n"
21
-
22
- prompt += f"User: {message}\nAI:"
23
- return prompt
24
-
25
- def generate(message: str, history: list[str]) -> str:
26
- prompt = build_prompt(history, message)
27
-
28
- result = generator(
29
- prompt,
30
- max_new_tokens=500,
31
- do_sample=True
32
- # temperature=0.7
33
- )
34
-
35
- text = result[0]["generated_text"]
36
-
37
- # Extract only the newly generated text (after the prompt)
38
- if text.startswith(prompt):
39
- response = text[len(prompt):].strip()
40
- else:
41
- # Fallback to original logic with better error handling
42
- response = text.split("AI:")[-1].strip()
43
-
44
- return response
 
1
+ from transformers import AutoTokenizer, AutoModelForCausalLM
2
+ import torch
3
+ from typing import List, Tuple
4
 
5
+ # Use a very small model for testing
6
+ class UniversalChatModel:
7
+ def __init__(self, model_name: str):
8
+ self.model_name = model_name
9
+ print(f"Loading tokenizer for {model_name}...")
10
+ self.tokenizer = AutoTokenizer.from_pretrained(model_name)
11
+
12
+ # Set padding token
13
+ if self.tokenizer.pad_token is None:
14
+ self.tokenizer.pad_token = self.tokenizer.eos_token or "<|endoftext|>"
15
+
16
+ print(f"Loading model {model_name}...")
17
+ self.model = AutoModelForCausalLM.from_pretrained(
18
+ model_name,
19
+ torch_dtype=torch.float16,
20
+ device_map="auto"
21
+ )
22
+
23
+ print("Model loaded successfully!")
24
+
25
+ def format_prompt_fallback(self, messages: List[dict]) -> str:
26
+ """Universal ChatML format for models without chat templates"""
27
+ chatml = ""
28
+ for message in messages:
29
+ role = message["role"]
30
+ content = message["content"]
31
+ chatml += f"<|im_start|>{role}\n{content}<|im_end|>\n"
32
+ chatml += "<|im_start|>assistant\n"
33
+ return chatml
34
+
35
+ def build_messages(self, history: List[Tuple[str, str]], current_message: str, system_prompt: str = None) -> List[dict]:
36
+ """Build universal message format"""
37
+ messages = []
38
+
39
+ # Add system prompt
40
+ if system_prompt:
41
+ messages.append({"role": "system", "content": system_prompt})
42
+
43
+ # Add history
44
+ for user_msg, assistant_msg in history:
45
+ messages.append({"role": "user", "content": user_msg})
46
+ messages.append({"role": "assistant", "content": assistant_msg})
47
+
48
+ # Add current message
49
+ messages.append({"role": "user", "content": current_message})
50
+
51
+ return messages
52
+
53
+ def format_prompt(self, messages: List[dict]) -> str:
54
+ """Format prompt using model's chat template or fallback"""
55
+ # Try model's built-in chat template
56
+ if hasattr(self.tokenizer, 'chat_template') and self.tokenizer.chat_template:
57
+ try:
58
+ prompt = self.tokenizer.apply_chat_template(
59
+ messages,
60
+ tokenize=False,
61
+ add_generation_prompt=True
62
+ )
63
+ return prompt
64
+ except Exception as e:
65
+ print(f"Warning: Failed to use chat template: {e}")
66
+ pass
67
+
68
+ # Fallback to ChatML
69
+ return self.format_prompt_fallback(messages)
70
+
71
+ def extract_response(self, prompt: str, generated_text: str) -> str:
72
+ """Universal response extraction"""
73
+ if hasattr(self.tokenizer, 'chat_template') and self.tokenizer.chat_template:
74
+ # Model has chat template - extract after prompt
75
+ if generated_text.startswith(prompt):
76
+ response = generated_text[len(prompt):].strip()
77
+ # Remove any remaining template markers
78
+ response = response.replace("<|im_end|>", "").strip()
79
+ response = response.replace("</s>", "").strip()
80
+ return response
81
+ else:
82
+ # Fallback: find first assistant response
83
+ if "assistant" in generated_text.lower():
84
+ parts = generated_text.lower().split("assistant")
85
+ response = parts[-1].strip()
86
+ response = response.replace("<|im_end|>", "").strip()
87
+ response = response.replace("</s>", "").strip()
88
+ return response
89
+ return generated_text.strip()
90
+ else:
91
+ # ChatML fallback
92
+ if "<|im_start|>assistant\n" in generated_text:
93
+ parts = generated_text.split("<|im_start|>assistant\n")
94
+ response = parts[-1].replace("<|im_end|>", "").strip()
95
+ response = response.replace("</s>", "").strip()
96
+ return response
97
+ elif generated_text.startswith(prompt):
98
+ response = generated_text[len(prompt):].strip()
99
+ response = response.replace("<|im_end|>", "").strip()
100
+ response = response.replace("</s>", "").strip()
101
+ return response
102
+ else:
103
+ return generated_text.strip()
104
+
105
+ def generate(self, message: str, history: List[Tuple[str, str]] = None, system_prompt: str = None) -> str:
106
+ """Generate response using universal chat template system"""
107
+ if history is None:
108
+ history = []
109
+
110
+ # Build messages
111
+ messages = self.build_messages(history, message, system_prompt)
112
+
113
+ # Format prompt
114
+ prompt = self.format_prompt(messages)
115
+
116
+ print(f"\n\n----- PROMPT -----\n{prompt}\n-----------------\n\n")
117
+
118
+ # Tokenize
119
+ inputs = self.tokenizer(prompt, return_tensors="pt", padding=True)
120
+ # Move to model device
121
+ inputs = {k: v.to(self.model.device) for k, v in inputs.items()}
122
+
123
+ # Generate
124
+ generation_config = {
125
+ "max_new_tokens": 150,
126
+ "do_sample": True,
127
+ "temperature": 0.7,
128
+ "pad_token_id": self.tokenizer.eos_token_id,
129
+ "eos_token_id": self.tokenizer.eos_token_id,
130
+ }
131
+
132
+ with torch.no_grad():
133
+ outputs = self.model.generate(
134
+ **inputs,
135
+ **generation_config
136
+ )
137
+
138
+ # Decode
139
+ generated_text = self.tokenizer.decode(outputs[0], skip_special_tokens=True)
140
+
141
+ print(f"\n\n----- GENERATED -----\n{generated_text}\n-------------------\n\n")
142
+
143
+ # Extract response
144
+ response = self.extract_response(prompt, generated_text)
145
+
146
+ return response
147
 
148
+ # Initialize with a tiny model for testing
149
+ # MODEL_NAME = "HuggingFaceH4/tiny-random-LlamaForCausalLM"
150
+ MODEL_NAME = "google/functiongemma-270m-it"
151
+ SYSTEM_PROMPT = "You are a helpful assistant."
 
152
 
153
+ # Create global model instance
154
+ print("Creating model instance...")
155
+ chat_model = UniversalChatModel(MODEL_NAME)
 
156
 
157
+ def generate(message: str, history: List[Tuple[str, str]]) -> str:
158
+ """Generate response using universal chat model"""
159
+ return chat_model.generate(message, history, SYSTEM_PROMPT)
160
 
161
+ if __name__ == "__main__":
162
+ # Quick test
163
+ print("Testing generation...")
164
+ reply = generate("What is 2+2?", [])
165
+ print(f"Final response: {reply}")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
api.py CHANGED
@@ -4,8 +4,8 @@ from ai_engine import generate
4
 
5
  app = FastAPI(title="Hello World AI")
6
 
7
- # VERY simple in-memory store
8
- sessions: dict[str, list[str]] = {}
9
 
10
 
11
  class ChatRequest(BaseModel):
@@ -23,7 +23,7 @@ def chat(req: ChatRequest):
23
 
24
  reply = generate(req.message, history)
25
 
26
- history.append(f"User: {req.message}")
27
- history.append(f"AI: {reply}")
28
 
29
  return {"reply": reply}
 
4
 
5
  app = FastAPI(title="Hello World AI")
6
 
7
+ # VERY simple in-memory store - now storing tuples (user_msg, assistant_msg)
8
+ sessions: dict[str, list[tuple[str, str]]] = {}
9
 
10
 
11
  class ChatRequest(BaseModel):
 
23
 
24
  reply = generate(req.message, history)
25
 
26
+ # Store as tuple (user_msg, assistant_msg)
27
+ history.append((req.message, reply))
28
 
29
  return {"reply": reply}
requirements.txt CHANGED
@@ -1,3 +1,4 @@
 
1
  fastapi
2
  uvicorn
3
  transformers
 
1
+ accelerate
2
  fastapi
3
  uvicorn
4
  transformers