Shrijanagain commited on
Commit
7595337
·
verified ·
1 Parent(s): 685b17e

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +213 -36
app.py CHANGED
@@ -11,10 +11,109 @@ from transformers import AutoTokenizer, AutoModelForCausalLM
11
  MODEL_ID = os.getenv("MODEL_ID", "WeiboAI/VibeThinker-3B")
12
  DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
13
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
14
  print("=" * 60)
15
- print("X-RUDRA M2 (CHAT)")
16
  print("MODEL:", MODEL_ID)
17
  print("DEVICE:", DEVICE)
 
18
  print("=" * 60)
19
 
20
  # ============================================================
@@ -30,52 +129,70 @@ if tokenizer.pad_token is None:
30
  print("Loading model...")
31
  model = AutoModelForCausalLM.from_pretrained(
32
  MODEL_ID,
33
- dtype=torch.float16 if DEVICE == "cuda" else torch.float32, # FIXED
34
  device_map="auto",
35
  trust_remote_code=True,
36
  )
37
  model.eval()
38
  print("MODEL READY")
39
 
 
40
  # ============================================================
41
- # GENERATION FUNCTION (history = list of dicts)
42
  # ============================================================
43
 
44
- @spaces.GPU
45
- def generate_response(message, history, max_tokens, temperature):
46
  """
47
- history: list of dicts with 'role' and 'content' (e.g. [{"role":"user","content":"Hi"}])
48
- returns: (new_message, updated_history) new_message = '' to clear input.
 
 
49
  """
50
- # Ensure history is always a list
51
- if history is None:
52
- history = []
53
-
54
- # Append the new user message
55
- history.append({"role": "user", "content": message})
56
 
57
- # Build prompt using chat template (if available)
58
- prompt = None
59
  if hasattr(tokenizer, "apply_chat_template") and tokenizer.chat_template is not None:
 
 
60
  try:
61
  prompt = tokenizer.apply_chat_template(
62
- history,
63
  tokenize=False,
64
  add_generation_prompt=True
65
  )
 
66
  except Exception as e:
67
  print("Chat template failed, falling back to manual format:", e)
68
 
69
- # Fallback manual format
70
- if prompt is None:
71
- manual = ""
72
- for turn in history:
73
- if turn["role"] == "user":
74
- manual += f"User: {turn['content']}\n"
75
- elif turn["role"] == "assistant":
76
- manual += f"Assistant: {turn['content']}\n"
77
- manual += "Assistant:"
78
- prompt = manual
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
79
 
80
  # Tokenize
81
  inputs = tokenizer(
@@ -86,7 +203,6 @@ def generate_response(message, history, max_tokens, temperature):
86
  padding=True,
87
  )
88
  inputs = {k: v.to(model.device) for k, v in inputs.items()}
89
-
90
  input_len = inputs["input_ids"].shape[-1]
91
 
92
  with torch.no_grad():
@@ -106,29 +222,72 @@ def generate_response(message, history, max_tokens, temperature):
106
  new_tokens = outputs[0][input_len:]
107
  answer = tokenizer.decode(new_tokens, skip_special_tokens=True).strip()
108
 
109
- # Append assistant response to history
 
110
  history.append({"role": "assistant", "content": answer})
111
 
112
- # Return empty string to clear the input box and the updated history
113
  return "", history
114
 
 
115
  # ============================================================
116
- # UI Chatbot WITHOUT 'type' argument (defaults to "messages")
117
  # ============================================================
118
 
119
- with gr.Blocks(title="X-RUDRA M2") as demo:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
120
  gr.Markdown(
121
  f"""
122
- # ⚡ X-RUDRA M2 – Chat
123
  **Model:** `{MODEL_ID}`
124
- **Device:** `{DEVICE}`
 
125
  """
126
  )
127
 
128
- # DO NOT add 'type=' here – let Gradio use its default.
129
- # The default on this version is "messages" (list of dicts).
130
  chatbot = gr.Chatbot(height=600, label="Conversation")
131
-
132
  with gr.Row():
133
  msg = gr.Textbox(placeholder="Ask anything...", scale=8)
134
  send = gr.Button("Send", variant="primary", scale=1)
@@ -148,6 +307,24 @@ with gr.Blocks(title="X-RUDRA M2") as demo:
148
  outputs=[msg, chatbot]
149
  )
150
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
151
  # ============================================================
152
  # START
153
  # ============================================================
 
11
  MODEL_ID = os.getenv("MODEL_ID", "WeiboAI/VibeThinker-3B")
12
  DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
13
 
14
+ # System prompt – defines the assistant's personality and constraints
15
+ SYSTEM_PROMPT = (
16
+ "You are X-RUDRA, a helpful, knowledgeable, and concise AI assistant. "
17
+ "# SEARCH DECISION SYSTEM PROMPT
18
+
19
+ ACT ONLY WHEN REQUIRED.
20
+
21
+ SEARCH WHEN:
22
+
23
+ * SEARCH
24
+ * BROWSE
25
+ * LOOKUP
26
+ * VERIFY
27
+ * CHECK
28
+ * FIND
29
+ * RESEARCH
30
+ * COMPARE CURRENT DATA
31
+ * CONFIRM LATEST DATA
32
+ * RETRIEVE EXTERNAL INFORMATION
33
+ * HANDLE UNCERTAIN FACTS
34
+ * HANDLE TIME-SENSITIVE INFORMATION
35
+ * HANDLE NICHE INFORMATION
36
+ * HANDLE LOCAL INFORMATION
37
+ * HANDLE CURRENT PRICES
38
+ * HANDLE CURRENT NEWS
39
+ * HANDLE CURRENT SPORTS
40
+ * HANDLE CURRENT PRODUCTS
41
+ * HANDLE CURRENT PEOPLE
42
+ * HANDLE CURRENT COMPANIES
43
+ * HANDLE CURRENT SOFTWARE
44
+ * HANDLE CURRENT DOCUMENTATION
45
+
46
+ DO NOT SEARCH WHEN:
47
+
48
+ * CHAT
49
+ * CONVERSE
50
+ * GREET
51
+ * JOKE
52
+ * BRAINSTORM
53
+ * EXPLAIN FROM KNOWN KNOWLEDGE
54
+ * REWRITE
55
+ * TRANSLATE
56
+ * SUMMARIZE PROVIDED TEXT
57
+ * WRITE
58
+ * CODE FROM PROVIDED REQUIREMENTS
59
+ * SOLVE SIMPLE REASONING
60
+ * ANSWER CASUAL QUESTIONS
61
+ * HANDLE TIMEPASS CONVERSATION
62
+
63
+ PRIORITIZE:
64
+
65
+ * USER INTENT
66
+ * ACCURACY
67
+ * FRESHNESS
68
+ * RELEVANCE
69
+ * PRIMARY SOURCES
70
+ * OFFICIAL SOURCES
71
+ * DIRECT EVIDENCE
72
+
73
+ AVOID:
74
+
75
+ * UNNECESSARY SEARCHES
76
+ * SEARCHING CASUAL CONVERSATION
77
+ * SEARCHING EVERY MESSAGE
78
+ * FABRICATING SEARCH RESULTS
79
+ * FABRICATING SOURCES
80
+ * FABRICATING CITATIONS
81
+ * USING OUTDATED INFORMATION WHEN FRESH INFORMATION IS REQUIRED
82
+
83
+ WHEN SEARCHING:
84
+
85
+ 1. IDENTIFY THE INFORMATION REQUIRED.
86
+ 2. FORMULATE PRECISE QUERIES.
87
+ 3. SEARCH RELEVANT SOURCES.
88
+ 4. VERIFY IMPORTANT CLAIMS.
89
+ 5. PREFER PRIMARY SOURCES.
90
+ 6. CROSS-CHECK CONFLICTING INFORMATION.
91
+ 7. DISTINGUISH FACT FROM INFERENCE.
92
+ 8. CITE SOURCES.
93
+ 9. ANSWER DIRECTLY.
94
+ 10. STOP SEARCHING WHEN SUFFICIENT EVIDENCE EXISTS.
95
+
96
+ WHEN NOT SEARCHING:
97
+
98
+ 1. UNDERSTAND THE REQUEST.
99
+ 2. USE AVAILABLE CONTEXT.
100
+ 3. ANSWER DIRECTLY.
101
+ 4. DO NOT PERFORM A SEARCH JUST TO APPEAR HELPFUL.
102
+
103
+ CORE RULE:
104
+
105
+ SEARCH FOR INFORMATION.
106
+ DO NOT SEARCH FOR CONVERSATION.
107
+
108
+ SEARCH ONLY WHEN SEARCHING IMPROVES ACCURACY, FRESHNESS, VERIFICATION, OR COMPLETENESS.
109
+ "
110
+ )
111
+
112
  print("=" * 60)
113
+ print("X-RUDRA M1 (CHAT + API)")
114
  print("MODEL:", MODEL_ID)
115
  print("DEVICE:", DEVICE)
116
+ print("SYSTEM PROMPT:", SYSTEM_PROMPT)
117
  print("=" * 60)
118
 
119
  # ============================================================
 
129
  print("Loading model...")
130
  model = AutoModelForCausalLM.from_pretrained(
131
  MODEL_ID,
132
+ dtype=torch.float16 if DEVICE == "cuda" else torch.float32,
133
  device_map="auto",
134
  trust_remote_code=True,
135
  )
136
  model.eval()
137
  print("MODEL READY")
138
 
139
+
140
  # ============================================================
141
+ # HELPER: BUILD PROMPT WITH SYSTEM + HISTORY
142
  # ============================================================
143
 
144
+ def build_prompt_with_system(history, new_user_message=None):
 
145
  """
146
+ Build a full prompt string from conversation history and an optional new user message.
147
+ history: list of dicts with 'role' and 'content' (user/assistant)
148
+ new_user_message: str (if provided, appended as user message)
149
+ Returns: prompt string ready for tokenization.
150
  """
151
+ # Create a copy of history and optionally add the new user message
152
+ messages = list(history) if history else []
153
+ if new_user_message is not None:
154
+ messages.append({"role": "user", "content": new_user_message})
 
 
155
 
156
+ # If the tokenizer has a chat template that supports system, use it
 
157
  if hasattr(tokenizer, "apply_chat_template") and tokenizer.chat_template is not None:
158
+ # Some templates expect a system message; we'll include it
159
+ full_messages = [{"role": "system", "content": SYSTEM_PROMPT}] + messages
160
  try:
161
  prompt = tokenizer.apply_chat_template(
162
+ full_messages,
163
  tokenize=False,
164
  add_generation_prompt=True
165
  )
166
+ return prompt
167
  except Exception as e:
168
  print("Chat template failed, falling back to manual format:", e)
169
 
170
+ # Fallback: manual formatting with system prompt
171
+ prompt = f"System: {SYSTEM_PROMPT}\n"
172
+ for turn in messages:
173
+ if turn["role"] == "user":
174
+ prompt += f"User: {turn['content']}\n"
175
+ elif turn["role"] == "assistant":
176
+ prompt += f"Assistant: {turn['content']}\n"
177
+ # Add a final "Assistant:" to prompt the model
178
+ prompt += "Assistant:"
179
+ return prompt
180
+
181
+
182
+ # ============================================================
183
+ # GENERATION FUNCTION (for chat UI)
184
+ # ============================================================
185
+
186
+ @spaces.GPU
187
+ def generate_response(message, history, max_tokens, temperature):
188
+ """
189
+ Takes the current message and history, returns updated history with assistant reply.
190
+ """
191
+ if history is None:
192
+ history = []
193
+
194
+ # Build prompt including the new user message
195
+ prompt = build_prompt_with_system(history, message)
196
 
197
  # Tokenize
198
  inputs = tokenizer(
 
203
  padding=True,
204
  )
205
  inputs = {k: v.to(model.device) for k, v in inputs.items()}
 
206
  input_len = inputs["input_ids"].shape[-1]
207
 
208
  with torch.no_grad():
 
222
  new_tokens = outputs[0][input_len:]
223
  answer = tokenizer.decode(new_tokens, skip_special_tokens=True).strip()
224
 
225
+ # Append user message and assistant response to history
226
+ history.append({"role": "user", "content": message})
227
  history.append({"role": "assistant", "content": answer})
228
 
 
229
  return "", history
230
 
231
+
232
  # ============================================================
233
+ # GENERATION FUNCTION (for API – standalone)
234
  # ============================================================
235
 
236
+ @spaces.GPU
237
+ def generate(prompt, max_tokens, temperature):
238
+ """
239
+ Standalone generation for API calls.
240
+ Expects a raw prompt string, returns the generated text.
241
+ """
242
+ # Build full prompt with system + user input
243
+ # We treat the input as a user message
244
+ messages = [{"role": "user", "content": prompt}]
245
+ full_prompt = build_prompt_with_system(messages)
246
+
247
+ inputs = tokenizer(
248
+ full_prompt,
249
+ return_tensors="pt",
250
+ truncation=True,
251
+ max_length=4096,
252
+ padding=True,
253
+ )
254
+ inputs = {k: v.to(model.device) for k, v in inputs.items()}
255
+ input_len = inputs["input_ids"].shape[-1]
256
+
257
+ with torch.no_grad():
258
+ outputs = model.generate(
259
+ **inputs,
260
+ max_new_tokens=int(max_tokens),
261
+ temperature=float(temperature),
262
+ do_sample=True,
263
+ top_p=0.95,
264
+ top_k=50,
265
+ repetition_penalty=1.15,
266
+ no_repeat_ngram_size=3,
267
+ pad_token_id=tokenizer.pad_token_id,
268
+ eos_token_id=tokenizer.eos_token_id,
269
+ )
270
+
271
+ new_tokens = outputs[0][input_len:]
272
+ answer = tokenizer.decode(new_tokens, skip_special_tokens=True).strip()
273
+ return answer
274
+
275
+
276
+ # ============================================================
277
+ # UI – Chat Interface
278
+ # ============================================================
279
+
280
+ with gr.Blocks(title="X-RUDRA M1") as demo:
281
  gr.Markdown(
282
  f"""
283
+ # ⚡ X-RUDRA M1 – Chat + API
284
  **Model:** `{MODEL_ID}`
285
+ **Device:** `{DEVICE}`
286
+ **System Prompt:** `{SYSTEM_PROMPT[:80]}...`
287
  """
288
  )
289
 
 
 
290
  chatbot = gr.Chatbot(height=600, label="Conversation")
 
291
  with gr.Row():
292
  msg = gr.Textbox(placeholder="Ask anything...", scale=8)
293
  send = gr.Button("Send", variant="primary", scale=1)
 
307
  outputs=[msg, chatbot]
308
  )
309
 
310
+ # ------------------------------------------------------------
311
+ # Hidden Interface for API – exposes /generate endpoint
312
+ # ------------------------------------------------------------
313
+ gr.Interface(
314
+ fn=generate,
315
+ inputs=[
316
+ gr.Textbox(label="prompt", lines=2),
317
+ gr.Slider(64, 2048, value=512, step=64, label="max_tokens"),
318
+ gr.Slider(0.1, 1.5, value=0.7, step=0.1, label="temperature")
319
+ ],
320
+ outputs=gr.Textbox(label="response"),
321
+ title="X-RUDRA M1 API",
322
+ description="Standalone generation endpoint.",
323
+ api_name="generate",
324
+ visible=False, # Hidden from UI, but API route is still active
325
+ )
326
+
327
+
328
  # ============================================================
329
  # START
330
  # ============================================================