S-Dreamer/CodeCraftLab
0
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 