LogiAI / backend /main.py
Logi6023's picture
Update backend/main.py
eb361c9 verified
Raw
History Blame Contribute Delete
1.79 kB
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
from typing import List, Optional
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM
from peft import PeftModel
app = FastAPI(title="LogiAI Backend")
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_methods=["*"],
allow_headers=["*"],
)
BASE_MODEL = "mistralai/Mistral-7B-Instruct-v0.3"
LORA_MODEL = "Logi6023/LogiAI"
print("🤖 LogiAI modell betöltése Hugging Face-ről...")
tokenizer = AutoTokenizer.from_pretrained(LORA_MODEL)
base_model = AutoModelForCausalLM.from_pretrained(
BASE_MODEL,
torch_dtype=torch.float16,
device_map="auto",
load_in_4bit=True,
)
model = PeftModel.from_pretrained(base_model, LORA_MODEL)
model.eval()
print("✅ LogiAI kész!")
TEMPLATE = """### Instruction:\n{}\n\n### Input:\n\n\n### Response:\n"""
class ChatRequest(BaseModel):
message: str
history: Optional[List[dict]] = []
class ChatResponse(BaseModel):
response: str
@app.get("/")
def root():
return {"status": "LogiAI fut!", "version": "1.0"}
@app.get("/health")
def health():
return {"status": "ok"}
@app.post("/chat", response_model=ChatResponse)
def chat(req: ChatRequest):
prompt = TEMPLATE.format(req.message)
inputs = tokenizer([prompt], return_tensors="pt").to(model.device)
with torch.no_grad():
outputs = model.generate(
**inputs,
max_new_tokens=512,
temperature=0.7,
do_sample=True,
pad_token_id=tokenizer.eos_token_id,
)
result = tokenizer.decode(outputs[0], skip_special_tokens=True)
response = result.split("### Response:")[-1].strip()
return ChatResponse(response=response)