Team Ai
Apppublic

daveyeb/bigcode-starcoder

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py93 linesDownload Raw Back to root
1import os2import torch3from fastapi import FastAPI, HTTPException4from pydantic import BaseModel5from transformers import AutoTokenizer, AutoModelForCausalLM6from typing import List, Optional7import uvicorn8 9# Initialize FastAPI app10app = FastAPI(title="StarCoder API")11 12 13# Define request model14class CodeGenerationRequest(BaseModel):15    prompt: str16    max_length: Optional[int] = 51217    temperature: Optional[float] = 0.718    top_p: Optional[float] = 0.9519    top_k: Optional[int] = 5020    num_return_sequences: Optional[int] = 121 22 23# Load model and tokenizer24@app.on_event("startup")25async def startup_event():26    global model, tokenizer27 28    # Use the correct model name - "bigcode/starcoder" should be "bigcode/starcoderbase"29 30    hf_token = os.getenv("HF_TOKEN")31    model_name = "bigcode/starcoder"  # This is the correct model name32 33    try:34        print("Loading tokenizer...")35        tokenizer = AutoTokenizer.from_pretrained(model_name, token=hf_token)36 37        print("Loading model...")38        model = AutoModelForCausalLM.from_pretrained(39            model_name, torch_dtype=torch.float16, device_map="auto", token=hf_token40        )41 42        print("Model and tokenizer loaded successfully")43    except Exception as e:44        print(f"Error loading model: {str(e)}")45        raise e46 47 48# Health check endpoint49@app.get("/")50async def health_check():51    return {"status": "healthy", "model": "bigcode/starcoderbase"}52 53 54# Code generation endpoint55@app.post("/generate")56async def generate_code(request: CodeGenerationRequest):57    try:58        inputs = tokenizer(request.prompt, return_tensors="pt").to(model.device)59 60        # Generate code61        with torch.no_grad():62            outputs = model.generate(63                inputs.input_ids,64                max_length=request.max_length,65                temperature=request.temperature,66                top_p=request.top_p,67                top_k=request.top_k,68                num_return_sequences=request.num_return_sequences,69                pad_token_id=tokenizer.eos_token_id,70            )71 72        # Decode the generated text73        generated_code = tokenizer.batch_decode(outputs, skip_special_tokens=True)74 75        return {76            "generated_code": generated_code,77            "parameters": {78                "prompt": request.prompt,79                "max_length": request.max_length,80                "temperature": request.temperature,81                "top_p": request.top_p,82                "top_k": request.top_k,83            },84        }85    except Exception as e:86        raise HTTPException(status_code=500, detail=str(e))87 88 89if __name__ == "__main__":90    # Get port from environment variable for Hugging Face Spaces91    port = int(os.environ.get("PORT", 7860))92    uvicorn.run("app:app", host="0.0.0.0", port=port)93