HachimiMT-demo / src /translator.py
ngocdang83's picture
feat(chunk): paragraph mode giu xuong dong - translator.py
f9f04c8 verified
Raw
History Blame
47.5 kB
"""HachimiMT Marian translation backend."""
from __future__ import annotations
import os
import time
from collections import OrderedDict
from concurrent.futures import Future, ThreadPoolExecutor
from dataclasses import dataclass
from enum import Enum
from functools import lru_cache
from pathlib import Path
from typing import Callable, Iterator
import sentencepiece as spm
from huggingface_hub import snapshot_download
from chunker import split_chunks
from hardware import (
HardwareProfile,
auto_all_gpus_by_default,
detect_hardware_profile,
resolve_gpu_indices,
)
from line_restore import assemble_paragraph_output
from token_chunker import source_token_ids, split_for_translation
import ctranslate2
ROOT = Path(__file__).resolve().parent.parent
MODELS_DIR = Path(os.environ.get("HACHIMIMT_MODELS_DIR", ROOT / "models"))
SPECIAL_ID_TO_TOKEN = {0: "<pad>", 1: "<s>", 2: "</s>", 3: "<unk>"}
SPECIAL_TOKEN_TO_ID = {token: token_id for token_id, token in SPECIAL_ID_TO_TOKEN.items()}
EOS_TOKEN_ID = 2
class Backend(str, Enum):
CT2 = "ct2"
TRANSFORMERS = "transformers"
@dataclass(frozen=True)
class ModelConfig:
label: str
model_id: str
use_marian_class: bool
generate_kwargs: dict
ct2_max_input_tokens: int
ct2_max_output_tokens: int
ct2_max_batch_size: int = 8
default_beam: int = 2
# Dung lượng xấp xỉ bản CT2 INT8 (MB) — chỉ để hiển thị badge "sẽ tải ~XMB".
# Khai khi thêm model mới; để None thì badge chỉ hiện "chưa tải" không kèm số.
ct2_size_mb: int | None = None
# Tên thư mục con chứa bản CT2 trên repo HF. Mặc định "ct2-int8_float32";
# một số repo dùng tên khác (vd "ct2-int8"), khai lại ở đây cho từng model.
ct2_subdir: str = "ct2-int8_float32"
# Nếu bản CT2 nằm ở repo khác với model gốc, khai riêng ở đây. PyTorch backend
# vẫn tải từ model_id, còn CT2 tải từ repo này vào cùng thư mục cache local.
ct2_model_id: str | None = None
MODELS: dict[str, ModelConfig] = {
"HachimiMT-60": ModelConfig(
label="HachimiMT-60",
model_id="ngocdang83/HachimiMT-60-zh-vi",
use_marian_class=True,
generate_kwargs={
"max_new_tokens": 300,
"no_repeat_ngram_size": 2,
"repetition_penalty": 1.2,
},
ct2_max_input_tokens=256,
ct2_max_output_tokens=300,
default_beam=2,
ct2_size_mb=57,
),
"HachimiMT-30": ModelConfig(
label="HachimiMT-30",
model_id="ngocdang83/HachimiMT-30-zh-vi",
use_marian_class=False,
generate_kwargs={
"max_length": 512,
},
# Input cap 160 (KHÔNG 512/256): model 37M scratch không giữ nổi context dài.
# Đo thực (sweep 512/256/160/câu trên đoạn dày entity): entity khớp 42%/62%/
# 94%/83%. 256 chunk-đầu vẫn VỠ nửa sau (lặp loạn). 160 = điểm ngọt: mỗi
# chunk ~165 token, đủ ngắn để không vỡ + đủ context giữ nhất quán tên riêng.
# Câu-thuần (context tối thiểu) NGƯỢC lại hại: 苍梧城→"Thương Ngô thành" (mất
# context → 城 dịch thành danh từ). Đường cong chữ U, 160 đáy. Output giữ 512.
ct2_max_input_tokens=160,
ct2_max_output_tokens=512,
default_beam=1,
ct2_size_mb=35,
),
"MoxhiMT-60": ModelConfig(
label="MoxhiMT-60",
model_id="DanVP/MoxhiMT-60",
use_marian_class=True,
generate_kwargs={
"max_new_tokens": 300,
"no_repeat_ngram_size": 2,
"repetition_penalty": 1.2,
},
ct2_max_input_tokens=256,
ct2_max_output_tokens=300,
default_beam=2,
ct2_size_mb=58,
ct2_subdir="ct2-int8", # repo này dùng tên thư mục CT2 khác
),
"MoxhiMT-30": ModelConfig(
label="MoxhiMT-30",
model_id="DanVP/MoxhiMT-30",
use_marian_class=True,
generate_kwargs={
"max_new_tokens": 300,
"no_repeat_ngram_size": 2,
"repetition_penalty": 1.2,
},
# 30M model drifts/hallucinates on dense paragraph chunks. Keep the
# source cap short enough to split long entity-heavy paragraphs.
ct2_max_input_tokens=160,
ct2_max_output_tokens=512,
default_beam=1,
ct2_size_mb=38,
),
"HirashibaMT-Medium": ModelConfig(
label="HirashibaMT-Medium",
model_id="Moleys/hirashiba-mt-medium",
use_marian_class=True,
generate_kwargs={
"max_new_tokens": 256,
},
ct2_max_input_tokens=128,
ct2_max_output_tokens=256,
default_beam=4,
ct2_size_mb=62,
ct2_model_id="ngungodan/hirashiba-mt-medium-ct2",
),
"HirashibaMT-Tiny": ModelConfig(
label="HirashibaMT-Tiny",
model_id="chi-vi/hirashiba-mt-tiny-zh-vi",
use_marian_class=True,
generate_kwargs={
"max_length": 512,
},
# Tiny uses a 4-layer Marian model. Keep paragraph chunks short and use
# greedy/low-beam decoding; higher beams can introduce light duplicates.
ct2_max_input_tokens=160,
ct2_max_output_tokens=512,
default_beam=1,
ct2_size_mb=17,
ct2_subdir="ct2-int8-keeppad",
ct2_model_id="ngungodan/hirashiba-mt-tiny-zh-vi-ct2",
),
}
# Model tải sẵn khi chạy setup (dùng được ngay); các model khác lazy-download.
DEFAULT_MODEL_KEY = "HachimiMT-60"
# Thư mục CT2 mặc định; model nào khác thì khai ModelConfig.ct2_subdir.
DEFAULT_CT2_SUBDIR = "ct2-int8_float32"
def _ct2_download_patterns(config: ModelConfig) -> list[str]:
return [
"config.json",
"generation_config.json",
"source.spm",
"target.spm",
"tokenizer.json",
"vocab.json",
"tokenizer_config.json",
f"{config.ct2_subdir}/*",
]
def _ct2_repo_id(config: ModelConfig) -> str:
return config.ct2_model_id or config.model_id
SourceTokenJobs = list[Future[list[list[str]]]]
def _env_int(name: str, default: int, *, min_value: int = 1, max_value: int = 1024) -> int:
raw = os.environ.get(name, "").strip()
if not raw:
return default
try:
return max(min_value, min(max_value, int(raw)))
except ValueError:
return default
def _batched(items: list[str], size: int) -> Iterator[list[str]]:
for start in range(0, len(items), size):
yield items[start : start + size]
def default_ct2_compute_type(device: str) -> str:
env_compute_type = os.environ.get("HACHIMIMT_COMPUTE_TYPE", "").strip()
if env_compute_type:
return env_compute_type
return "int8_float16" if device == "cuda" else "int8_float32"
def _ct2_gpu_index_attempts(gpu_indices: list[int]) -> list[list[int]]:
"""Try all requested GPUs first, then one GPU before giving up to CPU."""
if len(gpu_indices) <= 1:
return [list(gpu_indices)]
return [list(gpu_indices), [gpu_indices[0]]]
def _ct2_translator_kwargs(
*,
device: str,
compute_type: str,
intra_threads: int,
inter_threads: int,
gpu_indices: list[int] | None = None,
) -> tuple[dict[str, object], int, str | None]:
actual_inter_threads = max(1, int(inter_threads))
requested_intra_threads = max(1, int(intra_threads))
kwargs: dict[str, object] = dict(
device=device,
compute_type=compute_type,
inter_threads=actual_inter_threads,
)
worker_count = actual_inter_threads
device_indices_label = None
if device != "cuda":
kwargs["intra_threads"] = max(1, requested_intra_threads // worker_count)
return kwargs, worker_count, device_indices_label
if not gpu_indices:
raise RuntimeError("Không có GPU CUDA khả dụng.")
selected = list(dict.fromkeys(gpu_indices))
if len(selected) == 1:
# Luôn truyền device_index, kể cả single GPU: env "1" phải dùng GPU 1.
kwargs["device_index"] = selected[0]
else:
# CT2 inter_threads = replica trên MỖI device. Với nhiều GPU, giữ 1
# replica/GPU để tránh nhân VRAM và đã nhanh hơn trong benchmark T4x2.
kwargs["device_index"] = selected
kwargs["inter_threads"] = 1
actual_inter_threads = int(kwargs["inter_threads"])
worker_count = len(selected) * actual_inter_threads
kwargs["intra_threads"] = max(1, requested_intra_threads // worker_count)
device_indices_label = ",".join(str(i) for i in selected)
return kwargs, worker_count, device_indices_label
@lru_cache(maxsize=1)
def _optional_torch():
try:
import torch
except Exception:
return None
return torch
def _require_torch():
torch = _optional_torch()
if torch is None:
raise RuntimeError(
"Backend PyTorch cần cài torch. Engine mặc định CTranslate2 không cần torch. "
"Nếu muốn dùng PyTorch, cài torch rồi cài: pip install -r requirements-pytorch.txt"
)
return torch
def _torch_cuda_available() -> bool:
torch = _optional_torch()
if torch is None:
return False
try:
return bool(torch.cuda.is_available())
except Exception:
return False
def _torch_cuda_device_name() -> str | None:
torch = _optional_torch()
if torch is None:
return None
try:
if torch.cuda.is_available():
return str(torch.cuda.get_device_name(0))
except Exception:
return None
return None
def _torch_empty_cuda_cache() -> None:
torch = _optional_torch()
if torch is None:
return
try:
if torch.cuda.is_available():
torch.cuda.empty_cache()
except Exception:
return
class CT2SentencePieceTokenizer:
"""Minimal Marian SentencePiece tokenizer for CTranslate2 inference."""
pad_token_id = 0
def __init__(self, model_path: Path) -> None:
self._source_sp = spm.SentencePieceProcessor(model_file=str(model_path / "source.spm"))
self._target_sp = spm.SentencePieceProcessor(model_file=str(model_path / "target.spm"))
self._encode_cache_max = _env_int(
"HACHIMIMT_TOKEN_CACHE_SIZE",
0,
min_value=0,
max_value=500_000,
)
self._encode_cache: OrderedDict[str, list[int]] = OrderedDict()
self._cache_hits = 0
self._cache_misses = 0
def _full_encode_one(self, text: str) -> list[int]:
if self._encode_cache_max > 0:
cached = self._encode_cache.get(text)
if cached is not None:
self._encode_cache.move_to_end(text)
self._cache_hits += 1
return list(cached)
self._cache_misses += 1
token_ids = list(self._source_sp.encode(text, out_type=int))
token_ids.append(EOS_TOKEN_ID)
if self._encode_cache_max > 0:
self._encode_cache[text] = list(token_ids)
self._encode_cache.move_to_end(text)
while len(self._encode_cache) > self._encode_cache_max:
self._encode_cache.popitem(last=False)
return token_ids
def cache_stats(self) -> dict[str, int]:
return {
"token_cache_entries": len(self._encode_cache),
"token_cache_hits": self._cache_hits,
"token_cache_misses": self._cache_misses,
}
def _encode_one(
self,
text: str,
*,
truncation: bool = False,
max_length: int | None = None,
) -> list[int]:
token_ids = self._full_encode_one(text)
if truncation and max_length is not None and len(token_ids) > max_length:
token_ids = token_ids[:max_length]
if token_ids:
token_ids[-1] = EOS_TOKEN_ID
return token_ids
def __call__(
self,
text_or_texts: str | list[str],
*,
truncation: bool = False,
max_length: int | None = None,
padding: bool = False,
) -> dict[str, list[int] | list[list[int]]]:
del padding
if isinstance(text_or_texts, str):
return {
"input_ids": self._encode_one(
text_or_texts,
truncation=truncation,
max_length=max_length,
)
}
return {
"input_ids": [
self._encode_one(text, truncation=truncation, max_length=max_length)
for text in text_or_texts
]
}
def convert_ids_to_tokens(self, token_ids: list[int]) -> list[str]:
tokens: list[str] = []
for token_id in token_ids:
if token_id in SPECIAL_ID_TO_TOKEN:
tokens.append(SPECIAL_ID_TO_TOKEN[token_id])
else:
tokens.append(self._source_sp.id_to_piece(int(token_id)))
return tokens
def convert_tokens_to_ids(self, tokens: list[str]) -> list[int]:
token_ids: list[int] = []
for token in tokens:
if token in SPECIAL_TOKEN_TO_ID:
token_ids.append(SPECIAL_TOKEN_TO_ID[token])
else:
token_ids.append(int(self._target_sp.piece_to_id(token)))
return token_ids
def decode(self, token_ids: list[int], *, skip_special_tokens: bool = True) -> str:
return self.batch_decode([token_ids], skip_special_tokens=skip_special_tokens)[0]
def batch_decode(
self,
token_ids_batch: list[list[int]],
*,
skip_special_tokens: bool = True,
) -> list[str]:
decoded: list[str] = []
for token_ids in token_ids_batch:
if skip_special_tokens:
token_ids = [
token_id
for token_id in token_ids
if token_id not in SPECIAL_ID_TO_TOKEN
]
pieces = [self._target_sp.id_to_piece(int(token_id)) for token_id in token_ids]
decoded.append(self._target_sp.decode(pieces))
return decoded
def decode_tokens_batch(self, tokens_batch: list[list[str]]) -> list[str]:
decoded: list[str] = []
for tokens in tokens_batch:
pieces = [token for token in tokens if token not in SPECIAL_TOKEN_TO_ID]
decoded.append(self._target_sp.decode(pieces).strip())
return decoded
class CT2FastTokenizer:
"""Minimal tokenizer.json wrapper for CTranslate2 inference."""
def __init__(self, model_path: Path) -> None:
try:
from tokenizers import Tokenizer
except Exception as exc:
raise RuntimeError(
"Model này dùng tokenizer.json; cần cài package tokenizers."
) from exc
self._tokenizer = Tokenizer.from_file(str(model_path / "tokenizer.json"))
self.pad_token_id = self._tokenizer.token_to_id("<pad>")
self._eos_token_id = self._tokenizer.token_to_id("</s>")
self._encode_cache_max = _env_int(
"HACHIMIMT_TOKEN_CACHE_SIZE",
0,
min_value=0,
max_value=500_000,
)
self._encode_cache: OrderedDict[str, list[int]] = OrderedDict()
self._cache_hits = 0
self._cache_misses = 0
def _full_encode_one(self, text: str) -> list[int]:
if self._encode_cache_max > 0:
cached = self._encode_cache.get(text)
if cached is not None:
self._encode_cache.move_to_end(text)
self._cache_hits += 1
return list(cached)
self._cache_misses += 1
token_ids = list(self._tokenizer.encode(text).ids)
if self._encode_cache_max > 0:
self._encode_cache[text] = list(token_ids)
self._encode_cache.move_to_end(text)
while len(self._encode_cache) > self._encode_cache_max:
self._encode_cache.popitem(last=False)
return token_ids
def cache_stats(self) -> dict[str, int]:
return {
"token_cache_entries": len(self._encode_cache),
"token_cache_hits": self._cache_hits,
"token_cache_misses": self._cache_misses,
}
def _encode_one(
self,
text: str,
*,
truncation: bool = False,
max_length: int | None = None,
) -> list[int]:
token_ids = self._full_encode_one(text)
if truncation and max_length is not None and len(token_ids) > max_length:
token_ids = token_ids[:max_length]
if token_ids and self._eos_token_id is not None:
token_ids[-1] = self._eos_token_id
return token_ids
def __call__(
self,
text_or_texts: str | list[str],
*,
truncation: bool = False,
max_length: int | None = None,
padding: bool = False,
) -> dict[str, list[int] | list[list[int]]]:
del padding
if isinstance(text_or_texts, str):
return {
"input_ids": self._encode_one(
text_or_texts,
truncation=truncation,
max_length=max_length,
)
}
return {
"input_ids": [
self._encode_one(text, truncation=truncation, max_length=max_length)
for text in text_or_texts
]
}
def convert_ids_to_tokens(self, token_ids: list[int]) -> list[str]:
tokens: list[str] = []
for token_id in token_ids:
token = self._tokenizer.id_to_token(int(token_id))
if token is None:
raise ValueError(f"Token id không có trong vocab: {token_id}")
tokens.append(token)
return tokens
def convert_tokens_to_ids(self, tokens: list[str]) -> list[int]:
token_ids: list[int] = []
for token in tokens:
token_id = self._tokenizer.token_to_id(token)
if token_id is None:
raise ValueError(f"Token không có trong vocab: {token!r}")
token_ids.append(int(token_id))
return token_ids
def decode(self, token_ids: list[int], *, skip_special_tokens: bool = True) -> str:
return self.batch_decode([token_ids], skip_special_tokens=skip_special_tokens)[0]
def batch_decode(
self,
token_ids_batch: list[list[int]],
*,
skip_special_tokens: bool = True,
) -> list[str]:
return [
self._tokenizer.decode(token_ids, skip_special_tokens=skip_special_tokens)
for token_ids in token_ids_batch
]
def decode_tokens_batch(self, tokens_batch: list[list[str]]) -> list[str]:
token_ids_batch = [self.convert_tokens_to_ids(tokens) for tokens in tokens_batch]
return [
text.strip()
for text in self.batch_decode(token_ids_batch, skip_special_tokens=True)
]
def _load_ct2_tokenizer(model_path: Path):
if (model_path / "source.spm").exists() and (model_path / "target.spm").exists():
return CT2SentencePieceTokenizer(model_path)
if (model_path / "tokenizer.json").exists():
return CT2FastTokenizer(model_path)
raise RuntimeError("Không tìm thấy tokenizer CT2: cần source.spm/target.spm hoặc tokenizer.json.")
def model_local_dir(config: ModelConfig) -> Path:
return MODELS_DIR / config.model_id.split("/")[-1]
def _ct2_ready(path: Path, ct2_subdir: str = DEFAULT_CT2_SUBDIR) -> bool:
ct2_path = path / ct2_subdir
return ct2_path.is_dir() and any(ct2_path.iterdir())
def _pytorch_ready(path: Path) -> bool:
return any(path.glob("*.safetensors")) or any(path.glob("pytorch_model*.bin"))
def _tokenizer_ready(path: Path) -> bool:
has_sentencepiece = (path / "source.spm").exists() and (path / "target.spm").exists()
return has_sentencepiece or (path / "tokenizer.json").exists()
def is_model_downloaded(model_key: str, backend: Backend | str = Backend.CT2) -> bool:
"""Model (theo backend) đã có sẵn trong MODELS_DIR chưa — để UI hiện badge.
Dùng đúng điều kiện mà ensure_model_files() kiểm tra, nên kết quả khớp với
việc bấm Dịch có phải tải hay không.
"""
if isinstance(backend, str):
backend = Backend(backend)
if model_key not in MODELS:
return False
config = MODELS[model_key]
path = model_local_dir(config)
if backend == Backend.CT2:
weights_ready = _ct2_ready(path, config.ct2_subdir)
else:
weights_ready = _pytorch_ready(path)
return weights_ready and _tokenizer_ready(path)
def ensure_model_files(config: ModelConfig, backend: Backend) -> Path:
"""Download model vào MODELS_DIR nếu chưa có."""
local_dir = model_local_dir(config)
local_dir.mkdir(parents=True, exist_ok=True)
if backend == Backend.CT2:
if _ct2_ready(local_dir, config.ct2_subdir) and _tokenizer_ready(local_dir):
return local_dir
patterns = _ct2_download_patterns(config)
repo_id = _ct2_repo_id(config)
else:
if _pytorch_ready(local_dir) and _tokenizer_ready(local_dir):
return local_dir
patterns = None
repo_id = config.model_id
snapshot_download(
repo_id,
local_dir=str(local_dir),
allow_patterns=patterns,
)
return local_dir
def download_all_models(*, include_pytorch_weights: bool = False) -> list[Path]:
"""Tải trước tất cả model vào MODELS_DIR (dùng cho setup.bat).
Mặc định chỉ tải bản CT2 INT8 (~95 MB tổng) vì đó là engine mặc định.
PyTorch weights (~nặng gấp ~4 lần) chỉ tải khi include_pytorch_weights=True
hoặc tự động khi người dùng đổi sang engine PyTorch lần đầu (ensure_model_files).
"""
saved: list[Path] = []
for config in MODELS.values():
ensure_model_files(config, Backend.CT2)
saved.append(model_local_dir(config))
if include_pytorch_weights:
ensure_model_files(config, Backend.TRANSFORMERS)
return saved
class HachimiTranslator:
def __init__(self, profile: HardwareProfile | None = None) -> None:
self._profile = profile or detect_hardware_profile()
self._torch_device = "cuda" if _torch_cuda_available() else "cpu"
self._model_key: str | None = None
self._backend: Backend | None = None
self._tokenizer = None
self._torch_model = None
self._ct2_model = None
self._model_path: Path | None = None
self._ct2_threads = self._profile.ct2_threads
self._ct2_inter_threads = _env_int("HACHIMIMT_INTER_THREADS", 1, max_value=8)
self._ct2_window_multiplier = _env_int("HACHIMIMT_CT2_WINDOW_MULTIPLIER", 4, max_value=16)
self._tokenize_job_size = _env_int("HACHIMIMT_TOKENIZE_JOB_SIZE", 32, max_value=256)
batch_type = os.environ.get("HACHIMIMT_CT2_BATCH_TYPE", "tokens").strip().lower()
self._ct2_batch_type = batch_type if batch_type in {"examples", "tokens"} else "tokens"
self._ct2_compute_type: str | None = None
self._ct2_actual_intra_threads = self._ct2_threads
self._ct2_actual_inter_threads = self._ct2_inter_threads
self._ct2_worker_count = 1
self._ct2_device_indices_label: str | None = None
self._batch_size = self._profile.batch_size
self._tokenize_workers = self._profile.tokenize_workers
self._tokenize_pool: ThreadPoolExecutor | None = None
self._last_profile: dict[str, float | int | str] = {}
@property
def hardware_profile(self) -> HardwareProfile:
return self._profile
@property
def batch_size(self) -> int:
return self._batch_size
@property
def last_profile(self) -> dict[str, float | int | str]:
return dict(self._last_profile)
def _reset_profile(self) -> None:
self._last_profile = {}
def _profile_add(self, key: str, seconds: float) -> None:
self._last_profile[key] = float(self._last_profile.get(key, 0.0)) + seconds
def _profile_set(self, key: str, value: float | int | str) -> None:
self._last_profile[key] = value
def set_batch_size(self, batch_size: int) -> None:
self._batch_size = max(4, min(128, int(batch_size)))
def apply_hardware_profile(self, profile: HardwareProfile | None = None) -> None:
profile = profile or detect_hardware_profile()
threads_changed = profile.ct2_threads != self._ct2_threads
workers_changed = profile.tokenize_workers != self._tokenize_workers
self._profile = profile
self._ct2_threads = profile.ct2_threads
self._batch_size = profile.batch_size
self._tokenize_workers = profile.tokenize_workers
if workers_changed and self._tokenize_pool is not None:
self._tokenize_pool.shutdown(wait=False, cancel_futures=True)
self._tokenize_pool = None
if threads_changed and self._backend == Backend.CT2 and self._model_key:
model_key = self._model_key
self._unload_models()
self._load_ct2(MODELS[model_key])
self._model_key = model_key
self._backend = Backend.CT2
@property
def device(self) -> str:
if self._backend == Backend.CT2 and self._ct2_model is not None:
return self._ct2_model.device
return self._torch_device
def device_label(self) -> str:
"""Tên thiết bị inference thực tế (để phân biệt iGPU vs NVIDIA)."""
if self.device == "cuda":
return _torch_cuda_device_name() or self._profile.gpu_name or "CUDA GPU"
return "CPU"
@property
def backend(self) -> Backend | None:
return self._backend
def load(self, model_key: str, backend: Backend | str = Backend.CT2) -> str:
if isinstance(backend, str):
backend = Backend(backend)
if model_key not in MODELS:
raise ValueError(f"Unknown model: {model_key}")
if (
self._model_key == model_key
and self._backend == backend
and self._tokenizer is not None
and (self._ct2_model is not None or self._torch_model is not None)
):
return self._status_message(model_key, backend, cached=True)
config = MODELS[model_key]
self._unload_models()
if backend == Backend.CT2:
self._load_ct2(config)
else:
self._load_transformers(config)
self._model_key = model_key
self._backend = backend
# Warmup: dịch 1 câu để CT2/torch cấp phát buffer + nạp kernel. Cold-run
# đầu rất chậm nếu không warmup → lần dịch đầu của user nhanh ngay + số đo
# thời gian sạch (không gánh chi phí khởi tạo). Best-effort, không fatal.
try:
self.translate_chunk("你好。", beam_size=1)
except Exception: # noqa: BLE001 — warmup hỏng không được chặn load
pass
return self._status_message(model_key, backend, cached=False)
def _status_message(
self,
model_key: str,
backend: Backend,
*,
cached: bool,
beam_size: int | None = None,
) -> str:
prefix = "Model" if cached else "Đã tải"
config = MODELS[model_key]
engine = "CTranslate2 INT8" if backend == Backend.CT2 else "PyTorch"
msg = f"{prefix} {config.label} · {engine} · {self.device_label()}"
if backend == Backend.CT2 and self._ct2_compute_type:
msg += f" · compute={self._ct2_compute_type}"
window_multiplier = self._ct2_effective_window_multiplier()
window_part = f"window={self._ct2_window_multiplier}x"
if window_multiplier != self._ct2_window_multiplier:
window_part += f"/{window_multiplier}x"
msg += (
f" · batch_type={self._ct2_batch_type}"
f" · {window_part}"
f" · intra={self._ct2_actual_intra_threads}"
f" · inter={self._ct2_actual_inter_threads}"
)
if self._ct2_worker_count > 1:
msg += f" · workers={self._ct2_worker_count}"
if self._ct2_device_indices_label:
msg += f" · gpu={self._ct2_device_indices_label}"
if beam_size is not None:
msg += f" · beam={beam_size}"
return msg
@staticmethod
def clamp_beam(beam_size: int) -> int:
return max(1, min(4, int(beam_size)))
def _unload_models(self) -> None:
self._torch_model = None
self._ct2_model = None
self._tokenizer = None
self._model_path = None
self._ct2_compute_type = None
self._ct2_actual_intra_threads = self._ct2_threads
self._ct2_actual_inter_threads = self._ct2_inter_threads
self._ct2_worker_count = 1
self._ct2_device_indices_label = None
if self._tokenize_pool is not None:
self._tokenize_pool.shutdown(wait=False, cancel_futures=True)
self._tokenize_pool = None
if self._torch_device == "cuda":
_torch_empty_cuda_cache()
def _get_tokenize_pool(self) -> ThreadPoolExecutor:
if self._tokenize_pool is None:
self._tokenize_pool = ThreadPoolExecutor(
max_workers=self._tokenize_workers,
thread_name_prefix="hachimi-tokenize",
)
return self._tokenize_pool
def _tokenize_chunks_parallel(self, chunks: list[str]) -> list[list[str]]:
if not chunks:
return []
if len(chunks) <= self._tokenize_job_size or self._tokenize_workers <= 1:
return self._source_tokens_batch(chunks)
pool = self._get_tokenize_pool()
groups = list(_batched(chunks, self._tokenize_job_size))
nested = pool.map(self._source_tokens_batch, groups)
return [tokens for group in nested for tokens in group]
def _submit_tokenize_jobs(self, chunks: list[str]) -> SourceTokenJobs:
pool = self._get_tokenize_pool()
return [
pool.submit(self._source_tokens_batch, group)
for group in _batched(chunks, self._tokenize_job_size)
]
@staticmethod
def _collect_tokenize_jobs(jobs: SourceTokenJobs) -> list[list[str]]:
return [tokens for job in jobs for tokens in job.result()]
def _decode_ct2_results(self, results) -> list[str]:
start = time.perf_counter()
hypotheses = [result.hypotheses[0] for result in results]
decode_tokens_batch = getattr(self._tokenizer, "decode_tokens_batch", None)
if callable(decode_tokens_batch):
decoded = decode_tokens_batch(hypotheses)
self._profile_add("decode_s", time.perf_counter() - start)
return decoded
token_ids = [self._tokenizer.convert_tokens_to_ids(tokens) for tokens in hypotheses]
decoded = [
text.strip()
for text in self._tokenizer.batch_decode(token_ids, skip_special_tokens=True)
]
self._profile_add("decode_s", time.perf_counter() - start)
return decoded
def _load_ct2(self, config: ModelConfig) -> None:
model_path = ensure_model_files(config, Backend.CT2)
tokenizer = _load_ct2_tokenizer(model_path)
env_compute_type = os.environ.get("HACHIMIMT_COMPUTE_TYPE", "").strip()
ct2_device = "cuda" if self._profile.has_cuda else "cpu"
attempts: list[tuple[str, str, list[int] | None]] = []
if ct2_device == "cuda":
try:
cuda_count = ctranslate2.get_cuda_device_count()
except Exception:
cuda_count = 0
gpu_indices = resolve_gpu_indices(
cuda_count,
os.environ.get("HACHIMIMT_GPU_INDICES"),
auto_all=auto_all_gpus_by_default(),
)
compute_types = [default_ct2_compute_type("cuda")]
if not env_compute_type and "int8_float32" not in compute_types:
compute_types.append("int8_float32")
for compute_type in compute_types:
for candidate_indices in _ct2_gpu_index_attempts(gpu_indices):
attempts.append(("cuda", compute_type, candidate_indices))
attempts.append(("cpu", "int8_float32", None))
else:
cpu_compute_type = default_ct2_compute_type("cpu")
attempts.append(("cpu", cpu_compute_type, None))
if cpu_compute_type != "int8_float32":
attempts.append(("cpu", "int8_float32", None))
translator = None
last_error: Exception | None = None
for device, compute_type, gpu_indices in attempts:
try:
kwargs, worker_count, device_indices_label = _ct2_translator_kwargs(
device=device,
compute_type=compute_type,
intra_threads=self._ct2_threads,
inter_threads=self._ct2_inter_threads,
gpu_indices=gpu_indices,
)
translator = ctranslate2.Translator(
str(model_path / config.ct2_subdir), **kwargs
)
self._ct2_compute_type = compute_type
self._ct2_actual_intra_threads = int(kwargs["intra_threads"])
self._ct2_actual_inter_threads = int(kwargs["inter_threads"])
self._ct2_worker_count = worker_count
self._ct2_device_indices_label = device_indices_label
break
except Exception as exc:
last_error = exc
if translator is None:
raise RuntimeError("Không tải được CTranslate2 backend.") from last_error
self._tokenizer = tokenizer
self._ct2_model = translator
self._model_path = model_path
def _load_transformers(self, config: ModelConfig) -> None:
_require_torch()
try:
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer, MarianMTModel
except Exception as exc:
raise RuntimeError(
"Backend PyTorch cần transformers/sacremoses/safetensors. "
"Cài thêm: pip install -r requirements-pytorch.txt"
) from exc
model_path = ensure_model_files(config, Backend.TRANSFORMERS)
tokenizer = AutoTokenizer.from_pretrained(model_path)
if config.use_marian_class:
model = MarianMTModel.from_pretrained(model_path)
else:
model = AutoModelForSeq2SeqLM.from_pretrained(model_path)
model = model.to(self._torch_device).eval()
self._tokenizer = tokenizer
self._torch_model = model
def _chunk_text(self, text: str, chunk_mode: str) -> list[str]:
config = MODELS[self._model_key]
if self._backend == Backend.CT2 and self._tokenizer is not None:
return split_for_translation(
self._tokenizer,
text,
max_tokens=config.ct2_max_input_tokens,
chunk_mode=chunk_mode,
)
return split_chunks(text, mode=chunk_mode)
def _source_tokens(self, text: str) -> list[str]:
config = MODELS[self._model_key]
token_ids = source_token_ids(
self._tokenizer,
text,
max_length=config.ct2_max_input_tokens,
truncation=True,
)
return self._tokenizer.convert_ids_to_tokens(token_ids)
def _source_tokens_batch(self, chunks: list[str]) -> list[list[str]]:
config = MODELS[self._model_key]
encoded = self._tokenizer(
chunks,
truncation=True,
max_length=config.ct2_max_input_tokens,
padding=False,
)["input_ids"]
pad_id = self._tokenizer.pad_token_id
if pad_id is not None:
encoded = [
[token_id for token_id in token_ids if token_id != pad_id]
for token_ids in encoded
]
return [self._tokenizer.convert_ids_to_tokens(token_ids) for token_ids in encoded]
def _decode_tokens(self, tokens: list[str]) -> str:
token_ids = self._tokenizer.convert_tokens_to_ids(tokens)
return self._tokenizer.decode(token_ids, skip_special_tokens=True).strip()
def _torch_generate_kwargs(self, beam_size: int) -> dict:
config = MODELS[self._model_key]
kwargs = dict(config.generate_kwargs)
kwargs["num_beams"] = beam_size
if config.use_marian_class:
kwargs["early_stopping"] = beam_size > 1
return kwargs
def _runtime_batch_size(self, beam_size: int) -> int:
"""PyTorch tốn VRAM hơn theo beam — giảm batch để tránh OOM."""
if self._backend == Backend.CT2:
return self._batch_size
beam_size = self.clamp_beam(beam_size)
vram_factor = max(1, beam_size * 2)
return max(4, min(self._batch_size, 48 // vram_factor))
def _runtime_window_size(self, beam_size: int) -> int:
batch_size = self._runtime_batch_size(beam_size)
if self._backend == Backend.CT2:
return max(batch_size, batch_size * self._ct2_effective_window_multiplier())
return batch_size
def _ct2_effective_window_multiplier(self) -> int:
multiplier = self._ct2_window_multiplier
if self._ct2_batch_type == "tokens" and self._ct2_worker_count > 1:
# Multi-GPU needs enough queued chunks for CT2 to split into multiple
# token sub-batches; cap keeps large files from over-buffering.
multiplier *= min(self._ct2_worker_count * 2, 8)
return max(1, min(32, multiplier))
def _ct2_max_batch_size(self, config: ModelConfig) -> int:
if self._ct2_batch_type == "tokens":
return self._batch_size * config.ct2_max_input_tokens
return self._batch_size
def _translate_torch_batch(self, chunks: list[str], *, beam_size: int) -> list[str]:
if not chunks:
return []
torch = _require_torch()
config = MODELS[self._model_key]
max_length = 256 if config.use_marian_class else 512
inputs = self._tokenizer(
chunks,
return_tensors="pt",
padding=True,
truncation=True,
max_length=max_length,
).to(self._torch_device)
with torch.inference_mode():
outputs = self._torch_model.generate(
**inputs,
**self._torch_generate_kwargs(beam_size),
)
return [
self._tokenizer.decode(output, skip_special_tokens=True).strip()
for output in outputs
]
def translate_chunk(self, text: str, *, beam_size: int = 2) -> str:
if self._tokenizer is None:
raise RuntimeError("Chưa tải model. Gọi load() trước.")
beam_size = self.clamp_beam(beam_size)
if self._backend == Backend.CT2:
return self._translate_chunks_ct2([text], beam_size=beam_size)[0]
return self._translate_torch_batch([text], beam_size=beam_size)[0]
def count_chunks(self, text: str, chunk_mode: str = "sentence") -> int:
if not self._model_key:
raise RuntimeError("Chưa tải model. Gọi load() trước.")
return len(self._chunk_text(text, chunk_mode))
@staticmethod
def _ct2_repetition_kwargs(config: "ModelConfig") -> dict:
"""Lấy tham số chống lặp cho CT2 translate_batch TỪ generate_kwargs của
model (nhất quán với đường PyTorch). Chỉ lấy key CT2 hỗ trợ — model nào
không khai (vd HachimiMT-30) trả {} → không áp. Vá để đường CT2 (chạy
thật) cũng chống lặp như PyTorch; model có tiền sử lặp trên source ngắn."""
gk = config.generate_kwargs
kwargs = {}
if "no_repeat_ngram_size" in gk:
kwargs["no_repeat_ngram_size"] = gk["no_repeat_ngram_size"]
if "repetition_penalty" in gk:
kwargs["repetition_penalty"] = gk["repetition_penalty"]
return kwargs
def _translate_ct2_batch(
self,
chunks: list[str],
*,
beam_size: int,
source_batches: list[list[str]] | None = None,
) -> list[str]:
config = MODELS[self._model_key]
if source_batches is None:
tokenize_start = time.perf_counter()
source_batches = self._tokenize_chunks_parallel(chunks)
self._profile_add("tokenize_s", time.perf_counter() - tokenize_start)
infer_start = time.perf_counter()
results = self._ct2_model.translate_batch(
source_batches,
max_batch_size=self._ct2_max_batch_size(config),
batch_type=self._ct2_batch_type,
beam_size=beam_size,
max_decoding_length=config.ct2_max_output_tokens,
**self._ct2_repetition_kwargs(config),
)
self._profile_add("ct2_infer_s", time.perf_counter() - infer_start)
return self._decode_ct2_results(results)
def _translate_ct2_batch_pipelined(
self,
chunks: list[str],
*,
beam_size: int,
prefetched_tokens: SourceTokenJobs | None,
) -> list[str]:
if prefetched_tokens is not None:
wait_start = time.perf_counter()
source_batches = self._collect_tokenize_jobs(prefetched_tokens)
self._profile_add("tokenize_wait_s", time.perf_counter() - wait_start)
else:
source_batches = None
return self._translate_ct2_batch(
chunks,
beam_size=beam_size,
source_batches=source_batches,
)
def _translate_chunks_ct2(self, chunks: list[str], *, beam_size: int) -> list[str]:
if self._ct2_model is None:
raise RuntimeError("CTranslate2 chưa được tải.")
beam_size = self.clamp_beam(beam_size)
if not chunks:
return []
return self._translate_ct2_batch(chunks, beam_size=beam_size)
def translate_text_iter(
self,
text: str,
*,
chunk_mode: str = "sentence",
beam_size: int = 2,
) -> Iterator[tuple[int, int, str, list[tuple[int, str, str]] | None, str | None]]:
"""Yield (done, total, message, rows_or_none, full_text_or_none) sau mỗi batch."""
beam_size = self.clamp_beam(beam_size)
self._reset_profile()
self._profile_set("backend", self._backend.value if self._backend else "")
self._profile_set("beam", beam_size)
chunk_start = time.perf_counter()
chunks = self._chunk_text(text, chunk_mode)
self._profile_set("chunk_s", time.perf_counter() - chunk_start)
total = len(chunks)
self._profile_set("chunks", total)
yield 0, total, f"Đã chia {total} chunk, chuẩn bị dịch...", None, None
translations: list[str] = []
batch_size = self._runtime_batch_size(beam_size)
window_size = self._runtime_window_size(beam_size)
next_tokens: SourceTokenJobs | None = None
if self._backend == Backend.CT2 and total:
first_end = min(window_size, total)
submit_start = time.perf_counter()
next_tokens = self._submit_tokenize_jobs(chunks[:first_end])
self._profile_add("tokenize_submit_s", time.perf_counter() - submit_start)
for start in range(0, total, window_size):
end = min(start + window_size, total)
worker_label = (
f", workers {self._ct2_worker_count}"
if self._backend == Backend.CT2 and self._ct2_worker_count > 1
else ""
)
batch_label = (
f"window {window_size}, batch {batch_size}, {self._ct2_batch_type}{worker_label}"
if self._backend == Backend.CT2
else f"batch {batch_size}"
)
yield (
start,
total,
f"Đang dịch chunk {start + 1}{end}/{total} ({batch_label})...",
None,
None,
)
batch = chunks[start:end]
if self._backend == Backend.CT2:
current_tokens = next_tokens
next_tokens = None
next_start = start + window_size
next_end = min(next_start + window_size, total)
if next_start < total:
next_batch = chunks[next_start:next_end]
submit_start = time.perf_counter()
next_tokens = self._submit_tokenize_jobs(next_batch)
self._profile_add("tokenize_submit_s", time.perf_counter() - submit_start)
translations.extend(
self._translate_ct2_batch_pipelined(
batch,
beam_size=beam_size,
prefetched_tokens=current_tokens,
)
)
else:
translations.extend(self._translate_torch_batch(batch, beam_size=beam_size))
yield end, total, f"Đã xong {end}/{total} chunk", None, None
assemble_start = time.perf_counter()
if chunk_mode == "paragraph":
# Chunk gom nhiều dòng → model nuốt xuống dòng. Khôi phục bố cục dòng
# gốc (cho full_text) + ghép rows 1:1 dòng-gốc/dòng-dịch để bảng đối
# chiếu khớp bản tải về. Xem line_restore.assemble_paragraph_output.
rows, full_text = assemble_paragraph_output(text, chunks, translations)
else:
rows = [
(index, chunk, translated)
for index, (chunk, translated) in enumerate(zip(chunks, translations), start=1)
]
full_text = "\n".join(translations)
self._profile_add("assemble_s", time.perf_counter() - assemble_start)
cache_stats = getattr(self._tokenizer, "cache_stats", None)
if callable(cache_stats):
for key, value in cache_stats().items():
self._profile_set(key, value)
yield total, total, "Hoàn tất dịch.", rows, full_text
def translate_text(
self,
text: str,
*,
chunk_mode: str = "sentence",
beam_size: int = 2,
on_progress: Callable[[int, int, str], None] | None = None,
) -> tuple[list[tuple[int, str, str]], str]:
rows: list[tuple[int, str, str]] = []
full_text = ""
for done, total, message, result_rows, result_text in self.translate_text_iter(
text,
chunk_mode=chunk_mode,
beam_size=beam_size,
):
if on_progress:
on_progress(done, total, message)
if result_rows is not None and result_text is not None:
rows = result_rows
full_text = result_text
return rows, full_text