| |
| 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") |
|
|
| |
| |
| sbert_model_path = "" |
| sbert_model = SentenceTransformer(sbert_model_path) |
|
|
| |
| qwen_model_path = "fine_tuned_sbert_marketing" |
| qwen_tokenizer = AutoTokenizer.from_pretrained(qwen_model_path) |
| qwen_model = AutoModelForCausalLM.from_pretrained(qwen_model_path) |
|
|
| |
| class EmbeddingRequest(BaseModel): |
| sentences: list[str] |
|
|
| class ChatRequest(BaseModel): |
| prompt: str |
| max_new_tokens: int = 100 |
|
|
| |
| @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, |
| temperature=0.7, |
| top_p=0.9 |
| ) |
| response_text = qwen_tokenizer.decode(outputs[0], skip_special_tokens=True) |
| return {"response": response_text} |
|
|
| |
|
|