Team Ai
Apppublic

alter1/nova-llm-code

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
app.py129 linesDownload Raw Back to root
1import os2from typing import Dict, List, Optional, Union, Any3 4import torch5from fastapi import FastAPI, HTTPException, Request6from fastapi.middleware.cors import CORSMiddleware7from pydantic import BaseModel, Field8from transformers import AutoModelForCausalLM, AutoTokenizer, pipeline9 10# Initialize FastAPI app11app = FastAPI(title="NovaOS LLM Code Service (CodeLlama-7B)")12 13# Add CORS middleware14app.add_middleware(15    CORSMiddleware,16    allow_origins=["*"],17    allow_credentials=True,18    allow_methods=["*"],19    allow_headers=["*"],20)21 22# Model configuration23MODEL_ID = "codellama/CodeLlama-7b-hf"24MAX_LENGTH = 409625TEMPERATURE = 0.226TOP_P = 0.9527REPETITION_PENALTY = 1.228 29# Load model and tokenizer30@app.on_event("startup")31async def startup_event():32    global model, tokenizer, gen_pipeline33    34    print(f"Loading model: {MODEL_ID}")35    tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)36    model = AutoModelForCausalLM.from_pretrained(37        MODEL_ID, 38        torch_dtype=torch.float16,39        device_map="auto",40        load_in_8bit=True,  # Use 8-bit quantization for memory efficiency41        trust_remote_code=True42    )43    44    gen_pipeline = pipeline(45        "text-generation",46        model=model,47        tokenizer=tokenizer,48        max_length=MAX_LENGTH,49        temperature=TEMPERATURE,50        top_p=TOP_P,51        repetition_penalty=REPETITION_PENALTY,52        pad_token_id=tokenizer.eos_token_id53    )54    print("Model loaded successfully!")55 56# Define input/output models57class GenerationInput(BaseModel):58    prompt: str59    max_length: Optional[int] = Field(default=2048, description="Maximum length of the generated text")60    temperature: Optional[float] = Field(default=0.2, description="Sampling temperature")61    top_p: Optional[float] = Field(default=0.95, description="Top-p sampling")62    repetition_penalty: Optional[float] = Field(default=1.2, description="Repetition penalty")63    language: Optional[str] = Field(default=None, description="Target programming language")64 65class GenerationOutput(BaseModel):66    generated_text: str67    model_id: str68    usage: Dict[str, int]69 70# API endpoints71@app.get("/")72async def root():73    return {"message": "NovaOS LLM Code API is running", "model": MODEL_ID}74 75@app.post("/generate", response_model=GenerationOutput)76async def generate(input_data: GenerationInput):77    try:78        prompt = input_data.prompt79        80        # Add language-specific formatting if needed81        if input_data.language:82            prompt = f"Write code in {input_data.language} for the following: {prompt}"83        84        # Generate text85        output = gen_pipeline(86            prompt,87            max_length=min(input_data.max_length, MAX_LENGTH),88            temperature=input_data.temperature,89            top_p=input_data.top_p,90            repetition_penalty=input_data.repetition_penalty,91            return_full_text=False,92            num_return_sequences=1,93        )94        95        generated_text = output[0]["generated_text"]96        97        # Calculate token usage98        input_tokens = len(tokenizer.encode(prompt))99        output_tokens = len(tokenizer.encode(generated_text))100        101        return {102            "generated_text": generated_text,103            "model_id": MODEL_ID,104            "usage": {105                "input_tokens": input_tokens,106                "output_tokens": output_tokens,107                "total_tokens": input_tokens + output_tokens108            }109        }110    except Exception as e:111        raise HTTPException(status_code=500, detail=f"Generation failed: {str(e)}")112 113@app.post("/complete-code")114async def complete_code(input_data: GenerationInput):115    """Endpoint specifically optimized for code completion"""116    try:117        return await generate(input_data)118    except Exception as e:119        raise HTTPException(status_code=500, detail=f"Code completion failed: {str(e)}")120 121@app.get("/health")122async def health():123    return {"status": "healthy", "model": MODEL_ID}124 125# Run the API with uvicorn126if __name__ == "__main__":127    import uvicorn128    uvicorn.run(app, host="0.0.0.0", port=7860)129