import os import torch import spaces import gradio as gr from transformers import AutoTokenizer, AutoModelForCausalLM # ============================================================ # CONFIG # ============================================================ MODEL_ID = os.getenv("MODEL_ID", "WeiboAI/VibeThinker-3B") DEVICE = "cuda" if torch.cuda.is_available() else "cpu" # System prompt – use triple quotes for multi-line strings SYSTEM_PROMPT = """ You are X-RUDRA, a helpful, knowledgeable, and concise AI assistant. # SEARCH DECISION SYSTEM PROMPT ACT ONLY WHEN REQUIRED. SEARCH WHEN: * SEARCH * BROWSE * LOOKUP * VERIFY * CHECK * FIND * RESEARCH * COMPARE CURRENT DATA * CONFIRM LATEST DATA * RETRIEVE EXTERNAL INFORMATION * HANDLE UNCERTAIN FACTS * HANDLE TIME-SENSITIVE INFORMATION * HANDLE NICHE INFORMATION * HANDLE LOCAL INFORMATION * HANDLE CURRENT PRICES * HANDLE CURRENT NEWS * HANDLE CURRENT SPORTS * HANDLE CURRENT PRODUCTS * HANDLE CURRENT PEOPLE * HANDLE CURRENT COMPANIES * HANDLE CURRENT SOFTWARE * HANDLE CURRENT DOCUMENTATION DO NOT SEARCH WHEN: * CHAT * CONVERSE * GREET * JOKE * BRAINSTORM * EXPLAIN FROM KNOWN KNOWLEDGE * REWRITE * TRANSLATE * SUMMARIZE PROVIDED TEXT * WRITE * CODE FROM PROVIDED REQUIREMENTS * SOLVE SIMPLE REASONING * ANSWER CASUAL QUESTIONS * HANDLE TIMEPASS CONVERSATION PRIORITIZE: * USER INTENT * ACCURACY * FRESHNESS * RELEVANCE * PRIMARY SOURCES * OFFICIAL SOURCES * DIRECT EVIDENCE AVOID: * UNNECESSARY SEARCHES * SEARCHING CASUAL CONVERSATION * SEARCHING EVERY MESSAGE * FABRICATING SEARCH RESULTS * FABRICATING SOURCES * FABRICATING CITATIONS * USING OUTDATED INFORMATION WHEN FRESH INFORMATION IS REQUIRED WHEN SEARCHING: 1. IDENTIFY THE INFORMATION REQUIRED. 2. FORMULATE PRECISE QUERIES. 3. SEARCH RELEVANT SOURCES. 4. VERIFY IMPORTANT CLAIMS. 5. PREFER PRIMARY SOURCES. 6. CROSS-CHECK CONFLICTING INFORMATION. 7. DISTINGUISH FACT FROM INFERENCE. 8. CITE SOURCES. 9. ANSWER DIRECTLY. 10. STOP SEARCHING WHEN SUFFICIENT EVIDENCE EXISTS. WHEN NOT SEARCHING: 1. UNDERSTAND THE REQUEST. 2. USE AVAILABLE CONTEXT. 3. ANSWER DIRECTLY. 4. DO NOT PERFORM A SEARCH JUST TO APPEAR HELPFUL. CORE RULE: SEARCH FOR INFORMATION. DO NOT SEARCH FOR CONVERSATION. SEARCH ONLY WHEN SEARCHING IMPROVES ACCURACY, FRESHNESS, VERIFICATION, OR COMPLETENESS. """ print("=" * 60) print("X-RUDRA M1 (CHAT + API)") # Change to M2 for the other Space print("MODEL:", MODEL_ID) print("DEVICE:", DEVICE) print("SYSTEM PROMPT (first 100 chars):", SYSTEM_PROMPT[:100] + "...") print("=" * 60) # ============================================================ # LOAD MODEL # ============================================================ print("Loading tokenizer...") tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, trust_remote_code=True) if tokenizer.pad_token is None: tokenizer.pad_token = tokenizer.eos_token print("Loading model...") model = AutoModelForCausalLM.from_pretrained( MODEL_ID, dtype=torch.float16 if DEVICE == "cuda" else torch.float32, device_map="auto", trust_remote_code=True, ) model.eval() print("MODEL READY") # ============================================================ # HELPER: BUILD PROMPT WITH SYSTEM + HISTORY # ============================================================ def build_prompt_with_system(history, new_user_message=None): """ Build a full prompt string from conversation history and an optional new user message. history: list of dicts with 'role' and 'content' (user/assistant) new_user_message: str (if provided, appended as user message) Returns: prompt string ready for tokenization. """ messages = list(history) if history else [] if new_user_message is not None: messages.append({"role": "user", "content": new_user_message}) # If the tokenizer has a chat template that supports system, use it if hasattr(tokenizer, "apply_chat_template") and tokenizer.chat_template is not None: full_messages = [{"role": "system", "content": SYSTEM_PROMPT}] + messages try: prompt = tokenizer.apply_chat_template( full_messages, tokenize=False, add_generation_prompt=True ) return prompt except Exception as e: print("Chat template failed, falling back to manual format:", e) # Fallback: manual formatting with system prompt prompt = f"System: {SYSTEM_PROMPT}\n" for turn in messages: if turn["role"] == "user": prompt += f"User: {turn['content']}\n" elif turn["role"] == "assistant": prompt += f"Assistant: {turn['content']}\n" prompt += "Assistant:" return prompt # ============================================================ # GENERATION FUNCTION (for chat UI) # ============================================================ @spaces.GPU def generate_response(message, history, max_tokens, temperature): if history is None: history = [] prompt = build_prompt_with_system(history, message) inputs = tokenizer( prompt, return_tensors="pt", truncation=True, max_length=4096, padding=True, ) inputs = {k: v.to(model.device) for k, v in inputs.items()} input_len = inputs["input_ids"].shape[-1] with torch.no_grad(): outputs = model.generate( **inputs, max_new_tokens=int(max_tokens), temperature=float(temperature), do_sample=True, top_p=0.95, top_k=50, repetition_penalty=1.15, no_repeat_ngram_size=3, pad_token_id=tokenizer.pad_token_id, eos_token_id=tokenizer.eos_token_id, ) new_tokens = outputs[0][input_len:] answer = tokenizer.decode(new_tokens, skip_special_tokens=True).strip() history.append({"role": "user", "content": message}) history.append({"role": "assistant", "content": answer}) return "", history # ============================================================ # GENERATION FUNCTION (for API – standalone) # ============================================================ @spaces.GPU def generate(prompt, max_tokens, temperature): messages = [{"role": "user", "content": prompt}] full_prompt = build_prompt_with_system(messages) inputs = tokenizer( full_prompt, return_tensors="pt", truncation=True, max_length=4096, padding=True, ) inputs = {k: v.to(model.device) for k, v in inputs.items()} input_len = inputs["input_ids"].shape[-1] with torch.no_grad(): outputs = model.generate( **inputs, max_new_tokens=int(max_tokens), temperature=float(temperature), do_sample=True, top_p=0.95, top_k=50, repetition_penalty=1.15, no_repeat_ngram_size=3, pad_token_id=tokenizer.pad_token_id, eos_token_id=tokenizer.eos_token_id, ) new_tokens = outputs[0][input_len:] answer = tokenizer.decode(new_tokens, skip_special_tokens=True).strip() return answer # ============================================================ # UI – Chat Interface # ============================================================ with gr.Blocks(title="X-RUDRA M1") as demo: # Change to M2 for M2 Space gr.Markdown( f""" # ⚡ X-RUDRA M1 – Chat + API **Model:** `{MODEL_ID}` **Device:** `{DEVICE}` """ ) chatbot = gr.Chatbot(height=600, label="Conversation") with gr.Row(): msg = gr.Textbox(placeholder="Ask anything...", scale=8) send = gr.Button("Send", variant="primary", scale=1) with gr.Row(): max_tokens = gr.Slider(64, 2048, value=512, step=64, label="Max Tokens") temperature = gr.Slider(0.1, 1.5, value=0.7, step=0.1, label="Temperature") send.click( fn=generate_response, inputs=[msg, chatbot, max_tokens, temperature], outputs=[msg, chatbot] ) msg.submit( fn=generate_response, inputs=[msg, chatbot, max_tokens, temperature], outputs=[msg, chatbot] ) # Hidden API endpoint gr.Interface( fn=generate, inputs=[ gr.Textbox(label="prompt", lines=2), gr.Slider(64, 2048, value=512, step=64, label="max_tokens"), gr.Slider(0.1, 1.5, value=0.7, step=0.1, label="temperature") ], outputs=gr.Textbox(label="response"), title="X-RUDRA M1 API", description="Standalone generation endpoint.", api_name="generate", visible=False, ) # ============================================================ # START # ============================================================ if __name__ == "__main__": demo.launch(server_name="0.0.0.0", server_port=7860)