daveyeb/bigcode-starcoder
0
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 