"""HachimiMT Marian translation backend.""" from __future__ import annotations import os 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 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: "", 1: "", 2: "", 3: ""} 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" 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, }, ct2_max_input_tokens=512, ct2_max_output_tokens=512, default_beam=1, ct2_size_mb=38, ), } # 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", "source.spm", "target.spm", "vocab.json", "tokenizer_config.json", f"{config.ct2_subdir}/*", ] 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]: kwargs: dict[str, object] = dict( device=device, compute_type=compute_type, intra_threads=intra_threads, inter_threads=max(1, int(inter_threads)), ) if device != "cuda": return kwargs, max(1, int(inter_threads)), None 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"]) return kwargs, len(selected) * actual_inter_threads, ",".join(str(i) for i in selected) @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")) def _encode_one( self, text: str, *, truncation: bool = False, max_length: int | None = None, ) -> list[int]: token_ids = list(self._source_sp.encode(text, out_type=int)) token_ids.append(EOS_TOKEN_ID) 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 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: return (path / "source.spm").exists() or (path / "tokenizer_config.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) else: if _pytorch_ready(local_dir) and _tokenizer_ready(local_dir): return local_dir patterns = None snapshot_download( config.model_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_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 @property def hardware_profile(self) -> HardwareProfile: return self._profile @property def batch_size(self) -> int: return self._batch_size 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" · 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_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]: hypotheses = [result.hypotheses[0] for result in results] token_ids = [self._tokenizer.convert_tokens_to_ids(tokens) for tokens in hypotheses] return [ text.strip() for text in self._tokenizer.batch_decode(token_ids, skip_special_tokens=True) ] def _load_ct2(self, config: ModelConfig) -> None: model_path = ensure_model_files(config, Backend.CT2) tokenizer = CT2SentencePieceTokenizer(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_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: source_batches = self._tokenize_chunks_parallel(chunks) 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), ) 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: source_batches = self._collect_tokenize_jobs(prefetched_tokens) 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) chunks = self._chunk_text(text, chunk_mode) total = len(chunks) 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) next_tokens = self._submit_tokenize_jobs(chunks[:first_end]) 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] next_tokens = self._submit_tokenize_jobs(next_batch) 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 rows = [ (index, chunk, translated) for index, (chunk, translated) in enumerate(zip(chunks, translations), start=1) ] full_text = "\n".join(translations) 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