from fastapi import FastAPI from fastapi.responses import StreamingResponse from pydantic import BaseModel from transformers import ( AutoTokenizer, AutoModelForCausalLM, TextIteratorStreamer ) import transformers import torch from threading import Thread import asyncio import json import os from typing import Optional, List from huggingface_hub import login HF_TOKEN = os.getenv("HF_TOKEN") if HF_TOKEN: login(token=HF_TOKEN) print("✓ Hugging Face login successful") else: print("⚠ Warning: HF_TOKEN not found. Gated models may not work.") # ===================== # MODEL LOAD # ===================== #MODEL_NAME = "google/gemma-2-2b-it" MODEL_NAME = "google/gemma-3-1b-it" print("=" * 50) print("Transformers:", transformers.__version__) print("Loading model:", MODEL_NAME) tokenizer = AutoTokenizer.from_pretrained( MODEL_NAME, trust_remote_code=True, token=HF_TOKEN ) model = AutoModelForCausalLM.from_pretrained( MODEL_NAME, torch_dtype=torch.float32, low_cpu_mem_usage=True, trust_remote_code=True, token=HF_TOKEN ) print("Model loaded!") print("=" * 50) # ===================== # FASTAPI # ===================== app = FastAPI() # ===================== # REQUEST MODELS - সম্পূর্ণ কাস্টমাইজেবল # ===================== class ChatRequest(BaseModel): # messages: প্রতিটি মেসেজে role (system/user/assistant) এবং content থাকবে messages: List[dict] # জেনারেশন প্যারামিটার - সবগুলো ঐচ্ছিক temperature: Optional[float] = None max_tokens: Optional[int] = None top_p: Optional[float] = None top_k: Optional[int] = None repetition_penalty: Optional[float] = None do_sample: Optional[bool] = None use_cache: Optional[bool] = None # স্ট্রিমিং সক্রিয়/নিষ্ক্রিয় stream: Optional[bool] = False # ===================== # ROOT # ===================== @app.get("/") def root(): return { "status": "running", "model": MODEL_NAME, "transformers": transformers.__version__, "usage": { "description": "Send messages array with role and content", "example": { "messages": [ {"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "Hello!"} ], "temperature": 0.3, "max_tokens": 512, "stream": False } } } # ===================== # HEALTH # ===================== @app.get("/health") def health(): return {"status": "ok"} # ===================== # MAIN CHAT - সম্পূর্ণ কাস্টমাইজেবল # ===================== @app.post("/chat") async def chat(req: ChatRequest): try: # ===== ডিফল্ট প্যারামিটার ===== # ইউজার যদি প্যারামিটার না দেয়, তাহলে ডিফল্ট ব্যবহার হবে temperature = req.temperature if req.temperature is not None else 0.3 max_tokens = req.max_tokens if req.max_tokens is not None else 512 top_p = req.top_p if req.top_p is not None else 0.95 top_k = req.top_k if req.top_k is not None else 50 repetition_penalty = req.repetition_penalty if req.repetition_penalty is not None else 1.1 do_sample = req.do_sample if req.do_sample is not None else True use_cache = req.use_cache if req.use_cache is not None else True # ===== প্রম্পট তৈরি ===== # ইউজারের দেওয়া messages ব্যবহার করুন (সিস্টেম প্রম্পট সহ) prompt = tokenizer.apply_chat_template( req.messages, tokenize=False, add_generation_prompt=True ) # ===== টোকেনাইজ ===== inputs = tokenizer( prompt, return_tensors="pt" ) # ===== জেনারেট ===== with torch.no_grad(): outputs = model.generate( **inputs, max_new_tokens=max_tokens, temperature=temperature, top_p=top_p, top_k=top_k, repetition_penalty=repetition_penalty, do_sample=do_sample, use_cache=use_cache, pad_token_id=tokenizer.eos_token_id ) # ===== ডিকোড ===== response = tokenizer.decode( outputs[0][inputs.input_ids.shape[1]:], skip_special_tokens=True ) # ===== রেসপন্স ===== return { "response": response, "usage": { "prompt_tokens": inputs.input_ids.shape[1], "generated_tokens": len(outputs[0]) - inputs.input_ids.shape[1], "total_tokens": len(outputs[0]) }, "params": { "temperature": temperature, "max_tokens": max_tokens, "top_p": top_p, "top_k": top_k, "repetition_penalty": repetition_penalty, "do_sample": do_sample, "use_cache": use_cache } } except Exception as e: print(f"Error: {e}") import traceback traceback.print_exc() return { "error": str(e), "success": False } # ===================== # STREAM CHAT - সম্পূর্ণ কাস্টমাইজেবল # ===================== @app.post("/chat-stream") async def chat_stream(req: ChatRequest): try: # ===== ডিফল্ট প্যারামিটার ===== temperature = req.temperature if req.temperature is not None else 0.1 max_tokens = req.max_tokens if req.max_tokens is not None else 512 top_p = req.top_p if req.top_p is not None else 0.95 top_k = req.top_k if req.top_k is not None else 50 repetition_penalty = req.repetition_penalty if req.repetition_penalty is not None else 1.1 do_sample = req.do_sample if req.do_sample is not None else True use_cache = req.use_cache if req.use_cache is not None else True # ===== প্রম্পট তৈরি ===== prompt = tokenizer.apply_chat_template( req.messages, tokenize=False, add_generation_prompt=True ) # ===== টোকেনাইজ ===== inputs = tokenizer( prompt, return_tensors="pt" ) # ===== স্ট্রিমার ===== streamer = TextIteratorStreamer( tokenizer, skip_special_tokens=True, skip_prompt=True ) # ===== জেনারেশন কনফিগ ===== generation_kwargs = { **inputs, "max_new_tokens": max_tokens, "temperature": temperature, "top_p": top_p, "top_k": top_k, "repetition_penalty": repetition_penalty, "do_sample": do_sample, "use_cache": use_cache, "pad_token_id": tokenizer.eos_token_id, "streamer": streamer } # ===== থ্রেড ===== thread = Thread( target=model.generate, kwargs=generation_kwargs ) thread.start() # ===== জেনারেটর ===== async def generate(): full_response = "" for text in streamer: full_response += text yield f"data: {json.dumps({'token': text, 'full': full_response})}\n\n" await asyncio.sleep(0.01) # সম্পূর্ণ রেসপন্স শেষে তথ্য yield f"data: {json.dumps({'done': True, 'full_response': full_response})}\n\n" yield "data: [DONE]\n\n" # ===== রেসপন্স ===== return StreamingResponse( generate(), media_type="text/event-stream", headers={ "Cache-Control": "no-cache", "Connection": "keep-alive", "X-Accel-Buffering": "no" } ) except Exception as e: print(f"Stream Error: {e}") import traceback traceback.print_exc() return { "error": str(e), "success": False } # ===================== # MAIN # ===================== if __name__ == "__main__": import uvicorn uvicorn.run( app, host="0.0.0.0", port=7860 )