mitikoman / app.py
mhdacm's picture
Update app.py
c99f800 verified
Raw
History Blame Contribute Delete
2.95 kB
import os
os.environ["HF_HOME"] = "/app/cache/huggingface"
os.environ["TRANSFORMERS_CACHE"] = "/app/cache/huggingface"
from fastapi import FastAPI
from pydantic import BaseModel
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
import torch
import asyncio
from concurrent.futures import ThreadPoolExecutor
app = FastAPI()
# مدل NLLB-200 نسخه 600 میلیون پارامتری (سبک و مناسب برای Spaces رایگان)
print("⏳ در حال بارگذاری مدل ترجمه...")
MODEL_NAME = "Paulwalker4884/facebook-persian"
tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
model = AutoModelForSeq2SeqLM.from_pretrained(MODEL_NAME)
device = "cuda" if torch.cuda.is_available() else "cpu"
model = model.to(device)
print(f"✅ مدل بارگذاری شد. دستگاه: {device}")
# لیست کدهای زبان (فقط چند تای مهم)
SUPPORTED_LANGS = ["eng_Latn", "pes_Arab", "fas_Arab", "fra_Latn", "spa_Latn", "deu_Latn"]
class TranslateRequest(BaseModel):
text: str
src_lang: str = "eng_Latn"
tgt_lang: str = "pes_Arab"
executor = ThreadPoolExecutor(max_workers=1)
def translate_sync(text, src_lang, tgt_lang):
"""تابع همگام (synchronous) ترجمه"""
try:
# تنظیم زبان مبدأ برای tokenizer
tokenizer.src_lang = src_lang
# توکنایز کردن متن ورودی
inputs = tokenizer(text, return_tensors="pt", truncation=True, max_length=256)
inputs = {k: v.to(device) for k, v in inputs.items()}
# روش صحیح برای تنظیم زبان مقصد در نسخه‌های جدید transformers
# دریافت ID زبان مقصد با استفاده از تابع convert_tokens_to_ids
tgt_lang_id = tokenizer.convert_tokens_to_ids(tgt_lang)
# تولید ترجمه
generated_tokens = model.generate(
**inputs,
forced_bos_token_id=tgt_lang_id, # استفاده از ID محاسبه شده
max_length=256,
num_beams=4,
early_stopping=True
)
# دیکد کردن خروجی
translation = tokenizer.batch_decode(generated_tokens, skip_special_tokens=True)[0]
return translation
except Exception as e:
return f"خطا در ترجمه: {str(e)}"
@app.get("/")
def root():
return {
"message": "NLLB-200 Translator API - رایگان و آفلاین",
"supported_languages": SUPPORTED_LANGS,
"usage": "POST /translate با {'text': 'متن', 'src_lang': 'eng_Latn', 'tgt_lang': 'pes_Arab'}"
}
@app.post("/translate")
async def translate(request: TranslateRequest):
loop = asyncio.get_event_loop()
result = await loop.run_in_executor(
executor,
translate_sync,
request.text,
request.src_lang,
request.tgt_lang
)
return {"translation": result}