salim34's picture
Upload project folders
95d5743
Raw
History Blame Contribute Delete
1.54 kB
# app.py
from fastapi import FastAPI
from pydantic import BaseModel
from sentence_transformers import SentenceTransformer
from transformers import AutoTokenizer, AutoModelForCausalLM
import torch
app = FastAPI(title="SBERT + Qwen API")
# --------- Load Models ---------
# Path to your fine-tuned SBERT model
sbert_model_path = ""
sbert_model = SentenceTransformer(sbert_model_path)
# Path to your fine-tuned Qwen model
qwen_model_path = "fine_tuned_sbert_marketing"
qwen_tokenizer = AutoTokenizer.from_pretrained(qwen_model_path)
qwen_model = AutoModelForCausalLM.from_pretrained(qwen_model_path)
# --------- Pydantic request schemas ---------
class EmbeddingRequest(BaseModel):
sentences: list[str]
class ChatRequest(BaseModel):
prompt: str
max_new_tokens: int = 100
# --------- API Endpoints ---------
@app.post("/embed")
def get_embeddings(request: EmbeddingRequest):
embeddings = sbert_model.encode(request.sentences)
return {"embeddings": embeddings.tolist()}
@app.post("/chat")
def chat(request: ChatRequest):
inputs = qwen_tokenizer(request.prompt, return_tensors="pt")
outputs = qwen_model.generate(
**inputs,
max_new_tokens=request.max_new_tokens,
do_sample=True, # optional for randomness
temperature=0.7, # optional
top_p=0.9 # optional
)
response_text = qwen_tokenizer.decode(outputs[0], skip_special_tokens=True)
return {"response": response_text}
# --------- Run this app with: uvicorn app:app --reload ---------