Trans / app.py
Damonar's picture
Update app.py
8034bc8 verified
Raw History Blame Contribute Delete
12.6 kB
import os
import re
import time
import spaces
import torch
import gradio as gr
import pysubs2
from transformers import AutoTokenizer, AutoModelForCausalLM
# ── تنظیمات ──────────────────────────────────────────────
MODEL_ID = "google/translategemma-4b-it"
BATCH_SENTENCES = 8
MAX_NEW_TOKENS = 1536
CONTEXT_SENTENCES = 1
HF_TOKEN = os.environ.get("HF_TOKEN")
# ── بارگذاری مدل ─────────────────────────────────────────
tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, token=HF_TOKEN)
model = AutoModelForCausalLM.from_pretrained(
MODEL_ID,
torch_dtype=torch.bfloat16,
token=HF_TOKEN,
)
model.to("cuda")
model.eval()
# ── ابزارهای متن ──────────────────────────────────────────
def clean_subtitle_text(text):
text = re.sub(r"\\N", " ", text)
text = re.sub(r"\\n", " ", text)
text = re.sub(r"\{.*?\}", "", text)
text = re.sub(r"\s+", " ", text).strip()
return text
ENDERS = (".", "!", "?", "؟", "...", "…")
def build_sentence_groups(subs):
"""بلوک‌های SRT را بر اساس پایان جمله گروه‌بندی می‌کند — فقط برای
واحد ترجمه؛ خروجی نهایی بازهم به تعداد بلوک اصلی خواهد بود."""
groups, current = [], []
for i, sub in enumerate(subs):
t = clean_subtitle_text(sub.text)
if not t:
if current:
current.append(i)
continue
current.append(i)
if t.endswith(ENDERS) or i == len(subs) - 1:
groups.append(current)
current = []
if current:
groups.append(current)
return groups
def group_text(subs, group):
return clean_subtitle_text(" ".join(subs[i].text for i in group))
def compute_max_new_tokens(text):
word_count = max(1, len(text.split()))
return max(64, min(MAX_NEW_TOKENS, word_count * 10))
def parse_output(text, count):
text = (text or "").strip()
text = re.sub(r"```(?:text|plaintext|srt)?", "", text, flags=re.I).replace("```", "").strip()
text = re.sub(
r"(?im)^\s*\[?\s*SEGMENT\s+(\d+)\s*\]?\s*[::-]\s*",
lambda m: f"\n@@SEGMENT_{m.group(1)}@@\n",
text
)
parts = re.split(r"\n@@SEGMENT_(\d+)@@\n", "\n" + text + "\n")
found = {}
for i in range(1, len(parts), 2):
try:
n = int(parts[i])
value = parts[i + 1].strip()
if value:
found[n] = value
except Exception:
pass
if count == 1 and not found:
return [text]
if len(found) != count or any(i not in found for i in range(1, count + 1)):
raise ValueError(f"Expected {count} segments, got {len(found)}")
return [found[i] for i in range(1, count + 1)]
def build_prompt(sentences, count, context_before="", context_after=""):
context_block = ""
if context_before:
context_block += f"جمله‌ی قبل (فقط زمینه است؛ ترجمه‌اش نکن):\n{context_before}\n\n"
if context_after:
context_block += f"جمله‌ی بعد (فقط زمینه است؛ ترجمه‌اش نکن):\n{context_after}\n\n"
numbered = "\n\n".join(f"SEGMENT {n}:\n{s}" for n, s in enumerate(sentences, 1))
return f"""<start_of_turn>user
شما یک مترجم حرفه‌ای انگلیسی به فارسی هستید.
متن زیر شامل دقیقاً {count} جمله‌ی کامل و پیاپی از یک گفتگوی طبیعی است. هر جمله را با توجه به بافت کل گفتگو، به فارسی محاوره‌ای، خودمونی و روان ترجمه کن.
قوانین مهم:
- هر SEGMENT (جمله) را جداگانه، دقیقاً به همان ترتیب و به‌صورت کامل ترجمه کن.
- هرگز دو SEGMENT را ادغام نکن و هیچ‌کدام را نصف نکن.
- رویداد/فعلی که در «جمله‌ی قبل» گفته شده را دوباره تکرار نکن.
- کلمات عادی انگلیسی را فارسی کن؛ فقط اسم خاص/نرم‌افزار انگلیسی بماند.
- فقط خروجی خواسته‌شده (دقیقاً {count} SEGMENT) را بده.
خروجی را دقیقاً به این شکل بده:
SEGMENT 1: <ترجمه فارسی>
...
SEGMENT {count}: <ترجمه فارسی>
{context_block}جمله‌ها:
{numbered}<end_of_turn>
<start_of_turn>model
"""
@spaces.GPU(duration=90)
def run_translation(sentences, count, context_before="", context_after=""):
prompt = build_prompt(sentences, count, context_before, context_after)
inputs = tokenizer(prompt, return_tensors="pt")
inputs = {k: v.to(model.device) for k, v in inputs.items()}
joined = " ".join(sentences)
dynamic_max_tokens = compute_max_new_tokens(joined) + 32
with torch.inference_mode():
output = model.generate(
**inputs,
max_new_tokens=dynamic_max_tokens,
do_sample=False,
pad_token_id=tokenizer.eos_token_id,
)
generated = output[0][inputs["input_ids"].shape[-1]:]
answer = tokenizer.decode(generated, skip_special_tokens=True)
return parse_output(answer, count)
def translate_groups_retry(sentence_texts, start, end):
batch = sentence_texts[start:end]
count = len(batch)
context_before = sentence_texts[start - CONTEXT_SENTENCES] if start - CONTEXT_SENTENCES >= 0 else ""
context_after = sentence_texts[end] if end < len(sentence_texts) else ""
try:
return run_translation(batch, count, context_before, context_after)
except Exception as e:
print(f"⚠️ Batch failed: {e}")
print("🔁 Retrying sentence-by-sentence...")
result = []
for i in range(start, end):
try:
cb = sentence_texts[i - CONTEXT_SENTENCES] if i - CONTEXT_SENTENCES >= 0 else ""
ca = sentence_texts[i + 1] if i + 1 < len(sentence_texts) else ""
out = run_translation([sentence_texts[i]], 1, cb, ca)
result.append(out[0])
except Exception as e2:
print(f"❌ Sentence {i+1}: {e2}")
result.append(sentence_texts[i])
return result
# ── تقسیم هوشمند: چسبوندن مرزها به نزدیک‌ترین مکث طبیعی فارسی ──
BREAK_WORDS = {"و", "که", "اما", "ولی", "چون", "پس", "بعد", "یا", "تا"}
def _is_natural_break(words, pos):
"""آیا بریدن درست قبل از words[pos] طبیعیه؟ (بعد از ویرگول/نقطه، یا قبل از حرف ربط)"""
if pos <= 0 or pos >= len(words):
return False
prev = words[pos - 1]
curr = words[pos]
if prev.endswith(("،", ".", "!", "؟", "?")):
return True
if curr.strip("،.!؟?") in BREAK_WORDS:
return True
return False
def smart_split(translated_text, orig_word_lengths):
"""ترجمه‌ی روان یک جمله را دقیقاً به len(orig_word_lengths) تکه تقسیم می‌کند،
با چسباندن هر مرز به نزدیک‌ترین مکث طبیعی فارسی (نه فقط نسبت کلمه)."""
words = translated_text.split()
total_words = len(words)
n = len(orig_word_lengths)
if n == 1 or total_words == 0:
return [translated_text]
total_orig = sum(orig_word_lengths)
raw_cuts = []
acc = 0
for length in orig_word_lengths[:-1]:
acc += length
idx = round(total_words * acc / total_orig)
idx = max(1, min(total_words - 1, idx))
raw_cuts.append(idx)
window = max(2, total_words // 8)
used = set()
adjusted = []
for idx in raw_cuts:
best, best_dist = idx, None
for delta in range(0, window + 1):
for cand in {idx - delta, idx + delta}:
if 0 < cand < total_words and cand not in used and _is_natural_break(words, cand):
dist = abs(cand - idx)
if best_dist is None or dist < best_dist:
best_dist, best = dist, cand
if best_dist == 0:
break
adjusted.append(best)
used.add(best)
adjusted = sorted(set(adjusted))
if len(adjusted) != n - 1:
adjusted = raw_cuts # fallback به نسبت کلمه‌ای ساده اگر تعداد مرزها به‌هم خورد
pieces, prev = [], 0
for cut in adjusted:
pieces.append(" ".join(words[prev:cut]))
prev = cut
pieces.append(" ".join(words[prev:]))
return pieces
def fix_bidi(text):
LRI = "\u2066"
PDI = "\u2069"
pattern = re.compile(r'[A-Za-z0-9][A-Za-z0-9 _\-\.]*[A-Za-z0-9]|[A-Za-z0-9]')
return pattern.sub(lambda m: f"{LRI}{m.group(0)}{PDI}", text)
def translate_srt_file(input_path, progress=gr.Progress()):
subs = pysubs2.load(input_path, encoding="utf-8")
groups = build_sentence_groups(subs)
sentence_texts = [group_text(subs, g) for g in groups]
total = len(groups)
group_translations = [""] * total
started = time.time()
for start in range(0, total, BATCH_SENTENCES):
end = min(total, start + BATCH_SENTENCES)
group_translations[start:end] = translate_groups_retry(sentence_texts, start, end)
elapsed = time.time() - started
rate = end / elapsed if elapsed > 0 else 0
print(f"✅ {end}/{total} sentences | {rate:.2f}/sec", flush=True)
if progress:
progress(end / total, desc=f"ترجمه {end}/{total} جمله")
# ── تقسیم هوشمند هر جمله به تعداد بلوک‌های اصلی خودش ──
per_block_translation = [""] * len(subs)
for gi, group in enumerate(groups):
orig_lengths = [max(1, len(clean_subtitle_text(subs[i].text).split())) for i in group]
pieces = smart_split(group_translations[gi].strip(), orig_lengths)
for idx, piece in zip(group, pieces):
per_block_translation[idx] = piece
output = pysubs2.SSAFile()
preview = []
for i, sub in enumerate(subs):
original = sub.text.strip()
translated = fix_bidi(per_block_translation[i].strip()) if original else ""
output.append(pysubs2.SSAEvent(
start=sub.start,
end=sub.end,
text=translated
))
preview.append({"number": i + 1, "english": original, "persian": translated})
base = os.path.splitext(os.path.basename(input_path))[0]
output_path = f"/tmp/{base}_FA.srt"
output.save(output_path, encoding="utf-8")
print(f"✅ Finished in {(time.time()-started)/60:.1f} minutes")
return output_path, preview
def ui_translate(file, progress=gr.Progress()):
if file is None:
raise gr.Error("لطفاً یک فایل SRT انتخاب کنید.")
output_path, data = translate_srt_file(file, progress)
blocks = []
for x in data[:50]:
blocks.append(
f"### #{x['number']}\n\n"
f"🇬🇧 **English**\n\n{x['english']}\n\n"
f"🇮🇷 **فارسی**\n\n{x['persian']}\n\n---"
)
preview = "\n\n".join(blocks)
if len(data) > 50:
preview += f"\n\n> نمایش ۵۰ مورد اول از {len(data)} subtitle."
return output_path, preview
with gr.Blocks(title="TranslateGemma 4B SRT Translator") as demo:
gr.Markdown("""
# 🎬 English → Persian SRT Translator
### Powered by TranslateGemma 4B
هر جمله‌ی کامل یک‌جا و روان ترجمه می‌شود، سپس با چسباندن مرزها به نزدیک‌ترین مکث طبیعی فارسی،
دقیقاً به همان تعداد بلوک اصلی SRT تقسیم می‌شود — تعداد و زمان‌بندی بلوک‌ها بدون تغییر می‌ماند.
""")
with gr.Row():
input_file = gr.File(label="📂 English SRT", file_types=[".srt"], type="filepath")
output_file = gr.File(label="📥 Persian SRT")
translate_button = gr.Button("🚀 شروع ترجمه", variant="primary", size="lg")
gr.Markdown("## 👁️ پیش‌نمایش")
preview = gr.Markdown()
translate_button.click(
fn=ui_translate,
inputs=input_file,
outputs=[output_file, preview]
)
demo.queue(max_size=5).launch()