CoolFace
Apppublic

S-Dreamer/CodeCraftLab

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
model_utils.py226 linesDownload Raw Back to root
1import streamlit as st2import json3import os4from utils import add_log5 6# Initialize huggingface_models in session state if not present7if 'huggingface_models' not in st.session_state:8    st.session_state.huggingface_models = [9        "codegen-350M-mono",10        "codegen-2B-mono",11        "Salesforce/codegen-350M-mono",12        "Salesforce/codegen-2B-mono",13        "gpt2",14        "EleutherAI/gpt-neo-125M"15    ]16 17# Handle missing dependencies18try:19    import torch20    from transformers import AutoTokenizer, AutoModelForCausalLM21    TRANSFORMERS_AVAILABLE = True22except ImportError:23    TRANSFORMERS_AVAILABLE = False24    25    # Mock classes for demo purposes26    class DummyTokenizer:27        @classmethod28        def from_pretrained(cls, model_name):29            return cls()30            31        def __call__(self, text, **kwargs):32            return {"input_ids": [list(range(10))] * (1 if isinstance(text, str) else len(text))}33            34        def decode(self, token_ids, **kwargs):35            return "# Generated code placeholder\n\ndef example_function():\n    return 'Hello world!'"36        37        @property38        def eos_token(self):39            return "[EOS]"40            41        @property42        def eos_token_id(self):43            return 044            45        @property46        def pad_token(self):47            return None48            49        @pad_token.setter50        def pad_token(self, value):51            pass52            53    class DummyModel:54        @classmethod55        def from_pretrained(cls, model_name):56            return cls()57            58        def generate(self, input_ids, **kwargs):59            return [[1, 2, 3, 4, 5]]60            61        @property62        def config(self):63            class Config:64                @property65                def eos_token_id(self):66                    return 067                    68                @property69                def pad_token_id(self):70                    return 071                72                @pad_token_id.setter73                def pad_token_id(self, value):74                    pass75                    76            return Config()77            78    # Set aliases to match transformers79    AutoTokenizer = DummyTokenizer80    AutoModelForCausalLM = DummyModel81 82def list_available_huggingface_models():83    """84    List available code generation models from Hugging Face.85    86    Returns:87        list: List of model names88    """89    # Return the list stored in session state90    return st.session_state.huggingface_models91 92def get_model_and_tokenizer(model_name):93    """94    Load model and tokenizer from Hugging Face Hub.95    96    Args:97        model_name: Name of the model to load98        99    Returns:100        tuple: (model, tokenizer) or (None, None) if loading fails101    """102    try:103        add_log(f"Loading model and tokenizer: {model_name}")104        tokenizer = AutoTokenizer.from_pretrained(model_name)105        model = AutoModelForCausalLM.from_pretrained(model_name)106        add_log(f"Model and tokenizer loaded successfully: {model_name}")107        return model, tokenizer108    except Exception as e:109        add_log(f"Error loading model {model_name}: {str(e)}", "ERROR")110        return None, None111 112def save_trained_model(model_id, model, tokenizer):113    """114    Save trained model information to session state.115    116    Args:117        model_id: Identifier for the model118        model: The trained model119        tokenizer: The model's tokenizer120        121    Returns:122        bool: Success status123    """124    try:125        # Store model information in session state126        from datetime import datetime127        st.session_state.trained_models[model_id] = {128            'model': model,129            'tokenizer': tokenizer,130            'info': {131                'id': model_id,132                'created_at': datetime.now().strftime("%Y-%m-%d %H:%M:%S")133            }134        }135        add_log(f"Model {model_id} saved to session state")136        return True137    except Exception as e:138        add_log(f"Error saving model {model_id}: {str(e)}", "ERROR")139        return False140 141def list_trained_models():142    """143    List all trained models in session state.144    145    Returns:146        list: List of model IDs147    """148    if 'trained_models' in st.session_state:149        return list(st.session_state.trained_models.keys())150    return []151 152def generate_code(model_id, prompt, max_length=100, temperature=0.7, top_p=0.9):153    """154    Generate code using a trained model.155    156    Args:157        model_id: ID of the model to use158        prompt: Input prompt for code generation159        max_length: Maximum length of generated text160        temperature: Sampling temperature161        top_p: Nucleus sampling probability162        163    Returns:164        str: Generated code or error message165    """166    try:167        if model_id not in st.session_state.trained_models:168            return "Error: Model not found. Please select a valid model."169        170        model_data = st.session_state.trained_models[model_id]171        model = model_data['model']172        tokenizer = model_data['tokenizer']173        174        if TRANSFORMERS_AVAILABLE:175            # Tokenize the prompt176            inputs = tokenizer(prompt, return_tensors="pt", padding=True, truncation=True)177            178            # Generate text179            with torch.no_grad():180                outputs = model.generate(181                    inputs.input_ids,182                    max_length=max_length,183                    temperature=temperature,184                    top_p=top_p,185                    num_return_sequences=1,186                    pad_token_id=tokenizer.eos_token_id187                )188            189            # Decode the generated text190            generated_code = tokenizer.decode(outputs[0], skip_special_tokens=True)191        else:192            # Demo mode - return dummy generated code193            inputs = tokenizer(prompt)194            outputs = model.generate(inputs["input_ids"])195            generated_code = tokenizer.decode(outputs[0])196            197            # Add some context to the generated code based on the prompt198            if "fibonacci" in prompt.lower():199                generated_code = "def fibonacci(n):\n    if n <= 0:\n        return 0\n    elif n == 1:\n        return 1\n    else:\n        return fibonacci(n-1) + fibonacci(n-2)\n"200            elif "sort" in prompt.lower():201                generated_code = "def bubble_sort(arr):\n    n = len(arr)\n    for i in range(n):\n        for j in range(0, n-i-1):\n            if arr[j] > arr[j+1]:\n                arr[j], arr[j+1] = arr[j+1], arr[j]\n    return arr\n"202        203        # If the prompt is included in the output, remove it to get only the generated code204        if generated_code.startswith(prompt):205            generated_code = generated_code[len(prompt):]206            207        return generated_code208        209    except Exception as e:210        add_log(f"Error generating code: {str(e)}", "ERROR")211        return f"Error generating code: {str(e)}"212 213def get_model_info(model_id):214    """215    Get information about a model.216    217    Args:218        model_id: ID of the model219        220    Returns:221        dict: Model information222    """223    if 'trained_models' in st.session_state and model_id in st.session_state.trained_models:224        return st.session_state.trained_models[model_id]['info']225    return None226