david167/question-generation-api
0
1import os2import logging3import torch4from transformers import AutoTokenizer, AutoModelForCausalLM5import gradio as gr6import json7import re8 9# Configure logging10logging.basicConfig(level=logging.INFO)11logger = logging.getLogger(__name__)12 13class ModelManager:14 def __init__(self):15 self.model = None16 self.tokenizer = None17 self.device = None18 self.model_loaded = False19 self.load_model()20 21 def load_model(self):22 """Load the model and tokenizer"""23 try:24 logger.info("Starting model loading...")25 26 # Check if CUDA is available27 if torch.cuda.is_available():28 torch.cuda.set_device(0)29 self.device = "cuda:0"30 else:31 self.device = "cpu"32 logger.info(f"Using device: {self.device}")33 34 if self.device == "cuda:0":35 logger.info(f"GPU: {torch.cuda.get_device_name()}")36 logger.info(f"VRAM Available: {torch.cuda.get_device_properties(0).total_memory / 1024**3:.2f} GB")37 38 # Get HF token from environment39 hf_token = os.getenv("HF_TOKEN")40 41 logger.info("Loading Llama-3.1-8B-Instruct model...")42 base_model_name = "meta-llama/Llama-3.1-8B-Instruct"43 44 self.tokenizer = AutoTokenizer.from_pretrained(45 base_model_name,46 use_fast=True,47 trust_remote_code=True,48 token=hf_token49 )50 51 self.model = AutoModelForCausalLM.from_pretrained(52 base_model_name,53 torch_dtype=torch.float16 if self.device == "cuda:0" else torch.float32,54 device_map="auto" if self.device == "cuda:0" else None,55 trust_remote_code=True,56 token=hf_token57 )58 59 # Set pad token60 if self.tokenizer.pad_token is None:61 self.tokenizer.pad_token = self.tokenizer.eos_token62 63 self.model_loaded = True64 logger.info("✅ Model loaded successfully!")65 66 except Exception as e:67 logger.error(f"❌ Error loading model: {str(e)}")68 self.model_loaded = False69 70def generate_response(prompt, temperature=0.8, model_manager=None):71 """SIMPLE, WORKING GENERATION"""72 if not model_manager or not model_manager.model_loaded:73 return "Model not loaded"74 75 try:76 # Detect request type77 is_cot_request = any(phrase in prompt.lower() for phrase in [78 "return exactly this json array",79 "chain of thinking", 80 "verbatim"81 ])82 83 # Get model context84 max_context = getattr(model_manager.model.config, "max_position_embeddings", 8192)85 logger.info(f"Model context: {max_context} tokens")86 87 # SIMPLE PROMPT88 if is_cot_request:89 system_msg = "Generate JSON training data exactly as requested."90 else:91 system_msg = "You are a helpful AI assistant."92 93 formatted_prompt = f"""<|begin_of_text|><|start_header_id|>system<|end_header_id|>94 95{system_msg}96 97<|eot_id|><|start_header_id|>user<|end_header_id|>98 99{prompt}100 101<|eot_id|><|start_header_id|>assistant<|end_header_id|>102 103"""104 105 # REASONABLE TOKEN LIMITS106 if is_cot_request:107 max_new_tokens = 2048 # Reasonable for JSON108 min_new_tokens = 300 # Ensure completion109 else:110 max_new_tokens = 1024111 min_new_tokens = 50112 113 max_input_tokens = max_context - max_new_tokens - 100114 logger.info(f"Tokens: Input≤{max_input_tokens}, Output={min_new_tokens}-{max_new_tokens}")115 116 # Tokenize117 inputs = model_manager.tokenizer(118 formatted_prompt,119 return_tensors="pt",120 truncation=True,121 max_length=max_input_tokens122 )123 124 # Move to device125 if model_manager.device == "cuda:0":126 inputs = {k: v.to(next(model_manager.model.parameters()).device) for k, v in inputs.items()}127 128 # SIMPLE GENERATION129 with torch.no_grad():130 outputs = model_manager.model.generate(131 **inputs,132 max_new_tokens=max_new_tokens,133 min_new_tokens=min_new_tokens,134 temperature=temperature,135 top_p=0.9,136 do_sample=True,137 pad_token_id=model_manager.tokenizer.eos_token_id,138 early_stopping=False,139 repetition_penalty=1.1140 )141 142 # Decode143 full_response = model_manager.tokenizer.decode(outputs[0], skip_special_tokens=True)144 145 # Extract response146 if "<|start_header_id|>assistant<|end_header_id|>" in full_response:147 response = full_response.split("<|start_header_id|>assistant<|end_header_id|>", 1)[-1].strip()148 else:149 response = full_response[len(formatted_prompt):].strip()150 151 # For CoT, try to extract JSON152 if is_cot_request and '[' in response and ']' in response:153 json_match = re.search(r'\[.*\]', response, re.DOTALL)154 if json_match:155 candidate = json_match.group(0)156 if '"user"' in candidate and '"assistant"' in candidate:157 response = candidate158 159 logger.info(f"Response: {len(response)} chars")160 return response.strip()161 162 except Exception as e:163 logger.error(f"Generation error: {e}")164 return f"Error: {e}"165 166# Initialize model167model_manager = ModelManager()168 169def respond(message, history, temperature, json_mode=None, template=None):170 """Main API function matching original interface"""171 try:172 response = generate_response(message, temperature, model_manager)173 174 # Return in original format175 return [[176 {"role": "user", "metadata": None, "content": message, "options": None},177 {"role": "assistant", "metadata": None, "content": response, "options": None}178 ], ""]179 except Exception as e:180 logger.error(f"API Error: {e}")181 return [[182 {"role": "user", "metadata": None, "content": message, "options": None},183 {"role": "assistant", "metadata": None, "content": f"Error: {e}", "options": None}184 ], ""]185 186# Create simple interface187demo = gr.Interface(188 fn=respond,189 inputs=[190 gr.Textbox(label="Message", lines=5),191 gr.State(value=[]),192 gr.Slider(minimum=0.1, maximum=1.0, value=0.8, step=0.1, label="Temperature"),193 gr.Textbox(label="JSON Mode", value="", visible=False),194 gr.Textbox(label="Template", value="", visible=False)195 ],196 outputs=[197 gr.JSON(label="Response"),198 gr.Textbox(label="Status", visible=False)199 ],200 title="Question Generation API - Simple & Working",201 api_name="respond"202)203 204if __name__ == "__main__":205 demo.launch(server_name="0.0.0.0", server_port=7860, share=False)