from fastapi import FastAPI from pydantic import BaseModel from transformers import pipeline # Define the Pydantic model for the query class Query(BaseModel): query: str # Initialize FastAPI app app = FastAPI() # Load the models gpt2_model = "gpt2" # GPT-2 model for general purposes codebert_model = "microsoft/codebert-base" # CodeBERT model for code-related tasks # Initialize pipelines for GPT-2 and CodeBERT gpt2_generator = pipeline("text-generation", model=gpt2_model) codebert_generator = pipeline("fill-mask", model=codebert_model) # Function to detect if the query is related to coding or not def is_coding_query(user_input: str) -> bool: coding_keywords = ['code', 'debug', 'test', 'error', 'compile', 'script', 'function', 'class', 'python', 'java', 'javascript'] return any(keyword in user_input.lower() for keyword in coding_keywords) # Root endpoint @app.get("/") def read_root(): return {"message": "FastAPI on Hugging Face Spaces"} # Get request for predictions @app.get("/predict/") async def predict(query: str): # Check if the query is code-related or general if is_coding_query(query): # Use CodeBERT for code-related queries print("Using CodeBERT for code-related query...") masked_query = f"{query} " response = codebert_generator(masked_query) generated_text = response[0]['sequence'] if response else "Sorry, unable to generate code." else: # Use GPT-2 for general queries print("Using GPT-2 for general query...") response = gpt2_generator(query, max_length=1024) generated_text = response[0]["generated_text"] return {"response": generated_text} # Post request for predictions @app.post("/predict/") async def predict_post(query: Query): # Check if the query is code-related or general if is_coding_query(query.query): # Use CodeBERT for code-related queries print("Using CodeBERT for code-related query...") masked_query = f"{query.query} " response = codebert_generator(masked_query) generated_text = response[0]['sequence'] if response else "Sorry, unable to generate code." else: # Use GPT-2 for general queries print("Using GPT-2 for general query...") response = gpt2_generator(query.query, max_length=1024) generated_text = response[0]["generated_text"] return {"response": generated_text} # Run the application if __name__ == "__main__": import uvicorn uvicorn.run(app, host="0.0.0.0", port=7860)