alter1/nova-llm-code
0
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 