CoolFace
Apppublic

david167/question-generation-api

sourceHugging Faceapache-2.0updated 1y agoView on Hugging Face
0likes
gradio_app_simple.py205 linesDownload Raw Back to root
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)