Spaces:
Running on Zero
Running on Zero
| import gradio as gr | |
| import spaces | |
| import torch | |
| from transformers import ( | |
| AutoTokenizer, | |
| AutoModelForCausalLM | |
| ) | |
| from peft import PeftModel | |
| # ============================================================ | |
| # CONFIGURATION | |
| # ============================================================ | |
| BASE_MODEL = "Qwen/Qwen2.5-0.5B-Instruct" | |
| ADAPTER_MODEL = "dd253B/DhanushAI-0.5B" | |
| # ============================================================ | |
| # GLOBAL MODEL | |
| # ============================================================ | |
| tokenizer = None | |
| model = None | |
| # ============================================================ | |
| # LOAD MODEL | |
| # ============================================================ | |
| def load_model(): | |
| global tokenizer | |
| global model | |
| if model is not None: | |
| return | |
| print("====================================") | |
| print("Loading DhanushAI...") | |
| print("====================================") | |
| # ----------------------------- | |
| # Tokenizer | |
| # ----------------------------- | |
| print("Loading tokenizer...") | |
| tokenizer = AutoTokenizer.from_pretrained( | |
| BASE_MODEL | |
| ) | |
| if tokenizer.pad_token is None: | |
| tokenizer.pad_token = tokenizer.eos_token | |
| # ----------------------------- | |
| # Base model | |
| # ----------------------------- | |
| print("Loading base model...") | |
| base_model = AutoModelForCausalLM.from_pretrained( | |
| BASE_MODEL, | |
| torch_dtype=torch.float16 | |
| ) | |
| # ----------------------------- | |
| # LoRA adapter | |
| # ----------------------------- | |
| print("Loading DhanushAI adapter...") | |
| model = PeftModel.from_pretrained( | |
| base_model, | |
| ADAPTER_MODEL | |
| ) | |
| # ----------------------------- | |
| # Move to GPU | |
| # ----------------------------- | |
| model = model.to("cuda") | |
| model.eval() | |
| print("====================================") | |
| print("DhanushAI loaded successfully!") | |
| print("====================================") | |
| # ============================================================ | |
| # CHAT FUNCTION | |
| # ============================================================ | |
| def chat(message): | |
| # Load model after ZeroGPU allocation | |
| load_model() | |
| if message is None: | |
| return "Please enter a message." | |
| message = message.strip() | |
| if not message: | |
| return "Please enter a message." | |
| # ----------------------------- | |
| # Prompt | |
| # ----------------------------- | |
| prompt = f"""You are DhanushAI, a helpful AI assistant. | |
| User: {message} | |
| Assistant:""" | |
| # ----------------------------- | |
| # Tokenize | |
| # ----------------------------- | |
| inputs = tokenizer( | |
| prompt, | |
| return_tensors="pt" | |
| ) | |
| inputs = { | |
| key: value.to("cuda") | |
| for key, value in inputs.items() | |
| } | |
| # ----------------------------- | |
| # Generate | |
| # ----------------------------- | |
| with torch.no_grad(): | |
| outputs = model.generate( | |
| **inputs, | |
| max_new_tokens=200, | |
| temperature=0.7, | |
| top_p=0.9, | |
| do_sample=True, | |
| repetition_penalty=1.1 | |
| ) | |
| # ----------------------------- | |
| # Decode | |
| # ----------------------------- | |
| generated = tokenizer.decode( | |
| outputs[0], | |
| skip_special_tokens=True | |
| ) | |
| # ----------------------------- | |
| # Remove prompt | |
| # ----------------------------- | |
| if "Assistant:" in generated: | |
| answer = generated.split( | |
| "Assistant:", | |
| 1 | |
| )[1].strip() | |
| else: | |
| answer = generated.strip() | |
| return answer | |
| # ============================================================ | |
| # GRADIO UI + API | |
| # ============================================================ | |
| demo = gr.Interface( | |
| fn=chat, | |
| inputs=gr.Textbox( | |
| label="Message", | |
| placeholder="Ask DhanushAI something..." | |
| ), | |
| outputs=gr.Textbox( | |
| label="DhanushAI" | |
| ), | |
| title="DhanushAI", | |
| description="My custom AI model", | |
| api_name="chat" | |
| ) | |
| # ============================================================ | |
| # START | |
| # ============================================================ | |
| demo.launch() |