Files
python_script/flask-dev-api/blueprints/ai_translate.py
DelLevin-Home 05e733e872 22
2026-06-17 10:43:33 +08:00

622 lines
20 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""
ai-translate AI 翻译蓝图
支持 NLLB-200 本地翻译引擎
"""
import os
import re
import json
import uuid
import threading
import tempfile
from flask import Blueprint, render_template, request, jsonify, send_file
bp = Blueprint('ai_translate', __name__, url_prefix='/ai-translate')
try:
from config import BASE_DIR
except ImportError:
BASE_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
CONFIG_PATH = os.path.join(BASE_DIR, 'config', 'ai_translate_config.json')
# 引擎状态
_model = None
_model_lock = threading.Lock()
_engine_status = {'loaded': False, 'loading': False, 'engine_type': None, 'model_name': None, 'error': None}
_tasks = {}
SUBTITLE_EXTS = ('.srt', '.ass', '.ssa', '.vtt')
LANGUAGES = {
'zh': '中文', 'en': '英文', 'ja': '日文', 'ko': '韩文',
'fr': '法文', 'de': '德文', 'es': '西班牙文', 'ru': '俄文',
'th': '泰文', 'vi': '越南文', 'it': '意大利文',
'pt': '葡萄牙文', 'auto': '自动检测',
}
# NLLB-200 语言代码映射
_NLLB_LANG_MAP = {
'zh': 'zho_Hans', 'en': 'eng_Latn', 'ja': 'jpn_Jpan', 'ko': 'kor_Hang',
'fr': 'fra_Latn', 'de': 'deu_Latn', 'es': 'spa_Latn', 'ru': 'rus_Cyrl',
'th': 'tha_Thai', 'vi': 'vie_Latn', 'it': 'ita_Latn',
'pt': 'por_Latn',
}
# langdetect 返回值 → 我们的语言代码
_DETECT_LANG_MAP = {
'zh-cn': 'zh', 'zh-tw': 'zh', 'zh': 'zh',
'en': 'en', 'ja': 'ja', 'ko': 'ko',
'fr': 'fr', 'de': 'de', 'es': 'es', 'ru': 'ru',
'th': 'th', 'vi': 'vi', 'it': 'it', 'pt': 'pt',
}
ENGINE_TYPES = {
'nllb': {'name': 'NLLB-200', 'desc': 'HF 目录或模型 ID'},
}
def _default_config():
return {
'engine_type': 'nllb',
'nllb_dir': '',
'max_new_tokens': 512,
'batch_size': 8,
'last_nllb_model': '',
'last_src_lang': 'en',
'last_tgt_lang': 'zh',
}
def _load_config():
cfg = _default_config()
if os.path.exists(CONFIG_PATH):
try:
with open(CONFIG_PATH, 'r', encoding='utf-8') as f:
cfg.update(json.load(f))
except Exception:
pass
return cfg
def _save_config(cfg):
os.makedirs(os.path.dirname(CONFIG_PATH), exist_ok=True)
with open(CONFIG_PATH, 'w', encoding='utf-8') as f:
json.dump(cfg, f, ensure_ascii=False, indent=2)
# ==================== 字幕解析 ====================
def _parse_srt(text):
blocks = re.split(r'\n\s*\n', text.strip())
entries = []
for block in blocks:
lines = block.strip().split('\n')
if len(lines) < 3:
continue
try:
idx = int(lines[0].strip())
except ValueError:
continue
timecode = lines[1].strip()
content = '\n'.join(lines[2:]).strip()
if content:
entries.append((idx, timecode, content))
return entries
def _build_srt(entries):
parts = []
for idx, tc, content in entries:
parts.append(f'{idx}\n{tc}\n{content}')
return '\n\n'.join(parts) + '\n'
def _parse_vtt(text):
text = re.sub(r'^WEBVTT\s*\n', '', text.strip(), count=1)
blocks = re.split(r'\n\s*\n', text.strip())
entries = []
for i, block in enumerate(blocks):
lines = block.strip().split('\n')
if len(lines) < 2:
continue
timecode = None
content_lines = []
for line in lines:
if '-->' in line:
timecode = line.strip()
elif timecode is not None:
content_lines.append(line)
else:
if '-->' not in line:
continue
if timecode and content_lines:
entries.append((i + 1, timecode, '\n'.join(content_lines).strip()))
return entries
def _build_vtt(entries):
parts = ['WEBVTT\n']
for idx, tc, content in entries:
parts.append(f'{idx}\n{tc}\n{content}')
return '\n\n'.join(parts) + '\n'
def _parse_ass(text):
entries = []
for line in text.split('\n'):
line = line.strip()
if line.lower().startswith('dialogue:'):
entries.append(line)
return entries
def _ass_extract_text(dialogue_line):
parts = dialogue_line.split(',', 9)
if len(parts) < 10:
return dialogue_line, ''
text = parts[9]
text = re.sub(r'\{[^}]*\}', '', text)
text = text.replace('\\N', '\n').replace('\\n', '\n')
return dialogue_line, text.strip()
def _ass_replace_text(dialogue_line, new_text):
parts = dialogue_line.split(',', 9)
if len(parts) < 10:
return dialogue_line
translated = new_text.replace('\n', '\\N')
parts[9] = translated
return ','.join(parts)
def _build_ass(original_lines, translated_texts):
result = []
trans_idx = 0
for line in original_lines:
if line.lower().startswith('dialogue:'):
if trans_idx < len(translated_texts):
result.append(_ass_replace_text(line, translated_texts[trans_idx]))
trans_idx += 1
else:
result.append(line)
else:
result.append(line)
return '\n'.join(result)
def _detect_subtitle_format(filename, content):
ext = os.path.splitext(filename)[1].lower()
if ext == '.srt':
return 'srt'
if ext == '.vtt':
return 'vtt'
if ext in ('.ass', '.ssa'):
return 'ass'
if 'WEBVTT' in content[:20]:
return 'vtt'
if 'Dialogue:' in content:
return 'ass'
return 'srt'
# ==================== 模型加载与翻译 ====================
def _load_model(model_path, model_name):
global _model
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
import torch
_engine_status['loading'] = True
_engine_status['error'] = None
try:
device = 'cuda' if torch.cuda.is_available() else 'cpu'
tokenizer = AutoTokenizer.from_pretrained(model_path, clean_up_tokenization_spaces=True)
model = AutoModelForSeq2SeqLM.from_pretrained(model_path)
model = model.to(device)
model.eval()
_model = {'model': model, 'tokenizer': tokenizer, 'device': device}
_engine_status['loaded'] = True
_engine_status['model_name'] = model_name
_engine_status['engine_type'] = 'nllb'
_engine_status['error'] = None
return True, None
except Exception as e:
_engine_status['error'] = str(e)
return False, str(e)
finally:
_engine_status['loading'] = False
def _unload_model():
global _model
_model = None
_engine_status['loaded'] = False
_engine_status['model_name'] = None
_engine_status['engine_type'] = None
import gc
gc.collect()
try:
import torch
if torch.cuda.is_available():
torch.cuda.empty_cache()
except Exception:
pass
def _detect_lang(text):
"""用 langdetect 检测语言,返回我们的语言代码(如 'zh''en'"""
try:
from langdetect import detect
code = detect(text)
return _DETECT_LANG_MAP.get(code, _DETECT_LANG_MAP.get(code.split('-')[0], 'en'))
except Exception:
return 'en'
def _do_translate(texts, src_lang, tgt_lang, cfg):
if not _model:
return None, '模型未加载'
try:
import torch
model = _model['model']
tokenizer = _model['tokenizer']
device = _model['device']
nllb_tgt = _NLLB_LANG_MAP.get(tgt_lang, _NLLB_LANG_MAP.get('zh', 'zho_Hans'))
tgt_token_id = tokenizer.convert_tokens_to_ids(nllb_tgt)
# 自动检测时,用第一条文本检测语言
if src_lang == 'auto':
sample = ' '.join(texts[:3])
detected = _detect_lang(sample)
nllb_src = _NLLB_LANG_MAP.get(detected, 'eng_Latn')
else:
nllb_src = _NLLB_LANG_MAP.get(src_lang, 'eng_Latn')
results = []
for text in texts:
tokenizer.src_lang = nllb_src
inputs = tokenizer(text, return_tensors='pt', padding=True, truncation=True, max_length=512).to(device)
with torch.no_grad():
translated = model.generate(
**inputs,
forced_bos_token_id=tgt_token_id,
max_new_tokens=cfg.get('max_new_tokens', 512)
)
result = tokenizer.decode(translated[0], skip_special_tokens=True)
results.append(result)
return results, None
except Exception as e:
return None, str(e)
def _ensure_model(src_lang=None, tgt_lang=None):
if _engine_status['loaded'] and _model:
return True, None
cfg = _load_config()
models_dir = cfg.get('nllb_dir', '')
if not models_dir or not os.path.isdir(models_dir):
return False, '请在设置中配置 NLLB 模型目录'
# 用 last 模型或第一个可用
model_name = cfg.get('last_nllb_model', '')
if model_name:
full = os.path.join(models_dir, model_name)
if not (os.path.isdir(full) and os.path.exists(os.path.join(full, 'config.json'))):
model_name = ''
if not model_name:
for name in sorted(os.listdir(models_dir)):
full = os.path.join(models_dir, name)
if os.path.isdir(full) and os.path.exists(os.path.join(full, 'config.json')):
model_name = name
break
if not model_name:
return False, '未找到 NLLB 模型,请在设置中配置目录'
model_path = os.path.join(models_dir, model_name)
with _model_lock:
if _engine_status['loaded'] and _model:
return True, None
ok, load_err = _load_model(model_path, model_name)
if ok:
cfg['last_nllb_model'] = model_name
_save_config(cfg)
return ok, load_err
# ==================== 启动时自动加载模型 ====================
_auto_load_done = False
def _trigger_auto_load():
global _auto_load_done
if _auto_load_done:
return
_auto_load_done = True
threading.Thread(target=_ensure_model, daemon=True).start()
# ==================== 路由 ====================
@bp.route('/')
def page():
_trigger_auto_load()
cfg = _load_config()
return render_template('ai_translate.html', config=cfg, languages=LANGUAGES, engine_types=ENGINE_TYPES)
@bp.route('/model-status')
def model_status():
return jsonify({
'loaded': _engine_status['loaded'],
'loading': _engine_status['loading'],
'model_name': _engine_status.get('model_name'),
})
@bp.route('/config', methods=['GET'])
def get_config():
return jsonify(_load_config())
@bp.route('/config', methods=['POST'])
def save_config():
data = request.get_json()
cfg = _load_config()
for key in ('nllb_dir', 'last_src_lang', 'last_tgt_lang', 'max_new_tokens', 'batch_size'):
if key in data:
cfg[key] = data[key]
_save_config(cfg)
return jsonify({'success': True})
@bp.route('/models')
def list_models():
cfg = _load_config()
models_dir = cfg.get('nllb_dir', '')
models = []
if models_dir and os.path.isdir(models_dir):
for name in sorted(os.listdir(models_dir)):
full = os.path.join(models_dir, name)
if os.path.isdir(full) and os.path.exists(os.path.join(full, 'config.json')):
models.append({'name': name, 'type': 'huggingface', 'path': full})
return jsonify({'models': models, 'models_dir': models_dir, 'last_model': cfg.get('last_nllb_model', '')})
@bp.route('/languages')
def list_languages():
langs = [{'code': c, 'name': n} for c, n in LANGUAGES.items() if c != 'auto']
return jsonify({'src_langs': [{'code': 'auto', 'name': '自动检测'}] + langs, 'tgt_langs': langs})
@bp.route('/unload-model', methods=['POST'])
def unload_model():
with _model_lock:
_unload_model()
cfg = _load_config()
cfg['last_nllb_model'] = ''
_save_config(cfg)
return jsonify({'success': True, 'message': '模型已卸载'})
@bp.route('/load-model', methods=['POST'])
def load_model():
data = request.get_json() or {}
model_name = data.get('model_name', '').strip()
if not model_name:
return jsonify({'success': False, 'error': '请指定模型名称'}), 400
cfg = _load_config()
models_dir = cfg.get('nllb_dir', '')
if not models_dir or not os.path.isdir(models_dir):
return jsonify({'success': False, 'error': '请先配置 NLLB 模型目录'}), 400
model_path = os.path.join(models_dir, model_name)
if not os.path.isdir(model_path):
return jsonify({'success': False, 'error': f'模型目录不存在: {model_name}'}), 400
with _model_lock:
_unload_model()
ok, err = _load_model(model_path, model_name)
if ok:
cfg['last_nllb_model'] = model_name
_save_config(cfg)
return jsonify({'success': True, 'model_name': model_name})
return jsonify({'success': False, 'error': err or '加载失败'}), 400
@bp.route('/translate-text', methods=['POST'])
def translate_text():
data = request.get_json()
text = (data.get('text') or '').strip()
if not text:
return jsonify({'success': False, 'error': '请输入要翻译的文本'}), 400
src_lang = data.get('src_lang', 'en')
tgt_lang = data.get('tgt_lang', 'zh')
cfg = _load_config()
batch_size = cfg.get('batch_size', 8)
cfg['last_src_lang'] = src_lang
cfg['last_tgt_lang'] = tgt_lang
_save_config(cfg)
ok, err = _ensure_model(src_lang, tgt_lang)
if not ok:
return jsonify({'success': False, 'error': err}), 400
task_id = uuid.uuid4().hex
_tasks[task_id] = {'status': 'translating', 'progress': 0, 'result': None, 'error': None}
def _do_task():
try:
paragraphs = [p.strip() for p in text.split('\n') if p.strip()]
if not paragraphs:
paragraphs = [text]
all_results = []
total = len(paragraphs)
for i in range(0, total, batch_size):
batch = paragraphs[i:i + batch_size]
results, err = _do_translate(batch, src_lang, tgt_lang, cfg)
if err:
_tasks[task_id] = {'status': 'error', 'error': err}
return
all_results.extend(results)
_tasks[task_id]['progress'] = round((i + len(batch)) / total * 100)
_tasks[task_id] = {'status': 'done', 'result': '\n'.join(all_results), 'progress': 100}
except Exception as e:
_tasks[task_id] = {'status': 'error', 'error': str(e)}
threading.Thread(target=_do_task, daemon=True).start()
return jsonify({'success': True, 'task_id': task_id})
@bp.route('/translate-file', methods=['POST'])
def translate_file():
if 'subtitle' not in request.files:
return jsonify({'success': False, 'error': '请上传字幕文件'}), 400
f = request.files['subtitle']
if not f.filename:
return jsonify({'success': False, 'error': '请上传字幕文件'}), 400
ext = os.path.splitext(f.filename)[1].lower()
if ext not in SUBTITLE_EXTS:
return jsonify({'success': False, 'error': f'不支持的格式: {ext},支持 SRT/ASS/VTT'}), 400
src_lang = request.form.get('src_lang', 'en')
tgt_lang = request.form.get('tgt_lang', 'zh')
cfg = _load_config()
batch_size = cfg.get('batch_size', 8)
cfg['last_src_lang'] = src_lang
cfg['last_tgt_lang'] = tgt_lang
_save_config(cfg)
ok, err = _ensure_model(src_lang, tgt_lang)
if not ok:
return jsonify({'success': False, 'error': err}), 400
content = f.read().decode('utf-8', errors='replace')
fmt = _detect_subtitle_format(f.filename, content)
task_id = uuid.uuid4().hex
_tasks[task_id] = {'status': 'translating', 'progress': 0, 'result': None, 'error': None,
'filename': f.filename, 'format': fmt}
def _do_task():
try:
if fmt == 'srt':
entries = _parse_srt(content)
texts = [e[2] for e in entries]
elif fmt == 'vtt':
entries = _parse_vtt(content)
texts = [e[2] for e in entries]
elif fmt == 'ass':
full_lines = content.split('\n')
dialogue_lines = [l for l in full_lines if l.strip().lower().startswith('dialogue:')]
texts = []
for dl in dialogue_lines:
_, txt = _ass_extract_text(dl)
texts.append(txt if txt else '')
entries = dialogue_lines
else:
_tasks[task_id] = {'status': 'error', 'error': '未知字幕格式'}
return
if not texts:
_tasks[task_id] = {'status': 'error', 'error': '字幕文件中未找到可翻译的文本'}
return
all_results = []
total = len(texts)
for i in range(0, total, batch_size):
batch = texts[i:i + batch_size]
batch = [t for t in batch if t.strip()]
if not batch:
all_results.extend(texts[i:i + batch_size])
_tasks[task_id]['progress'] = round((i + batch_size) / total * 100)
continue
results, err = _do_translate(batch, src_lang, tgt_lang, cfg)
if err:
_tasks[task_id] = {'status': 'error', 'error': err}
return
ri = 0
for j in range(i, min(i + batch_size, total)):
if texts[j].strip():
all_results.append(results[ri] if ri < len(results) else texts[j])
ri += 1
else:
all_results.append('')
_tasks[task_id]['progress'] = round(min(i + batch_size, total) / total * 100)
if fmt == 'srt':
translated_entries = [(entries[i][0], entries[i][1], all_results[i]) for i in range(len(entries))]
output = _build_srt(translated_entries)
elif fmt == 'vtt':
translated_entries = [(entries[i][0], entries[i][1], all_results[i]) for i in range(len(entries))]
output = _build_vtt(translated_entries)
elif fmt == 'ass':
output = _build_ass(content.split('\n'), all_results)
else:
output = '\n'.join(all_results)
out_path = os.path.join(tempfile.gettempdir(), f'ai_trans_{task_id}{ext}')
with open(out_path, 'w', encoding='utf-8') as fout:
fout.write(output)
_tasks[task_id] = {
'status': 'done', 'progress': 100,
'result': output, 'file_path': out_path,
'filename': f.filename, 'format': fmt,
}
except Exception as e:
_tasks[task_id] = {'status': 'error', 'error': str(e)}
threading.Thread(target=_do_task, daemon=True).start()
return jsonify({'success': True, 'task_id': task_id})
@bp.route('/task-status/<task_id>')
def task_status(task_id):
task = _tasks.get(task_id)
if not task:
return jsonify({'success': False, 'error': '任务不存在'}), 404
return jsonify({'success': True, **task})
@bp.route('/download/<task_id>')
def download(task_id):
task = _tasks.get(task_id)
if not task or task.get('status') != 'done':
return jsonify({'error': '文件不存在'}), 404
file_path = task.get('file_path')
if not file_path or not os.path.exists(file_path):
return jsonify({'error': '文件不存在'}), 404
orig = task.get('filename', 'translated')
base, ext = os.path.splitext(orig)
return send_file(file_path, as_attachment=True, download_name=f'{base}_translated{ext}')
@bp.route('/cleanup/<task_id>', methods=['POST'])
def cleanup(task_id):
task = _tasks.pop(task_id, None)
if task and task.get('file_path'):
try:
os.remove(task['file_path'])
except OSError:
pass
return jsonify({'success': True})