QAi / app.py
Sampathl's picture
Update app.py
b7708fc verified
Raw
History Blame Contribute Delete
2.56 kB
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} <mask>"
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} <mask>"
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)