diff --git a/.gitattributes b/.gitattributes
index a6344aac8c09253b3b630fb776ae94478aa0275b..9f56080c719146e607ffbc69b28496e9e575d653 100644
--- a/.gitattributes
+++ b/.gitattributes
@@ -33,3 +33,72 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
*.zip filter=lfs diff=lfs merge=lfs -text
*.zst filter=lfs diff=lfs merge=lfs -text
*tfevents* filter=lfs diff=lfs merge=lfs -text
+MTP/gemma-4-12B-it-BF16-MTP.gguf filter=lfs diff=lfs merge=lfs -text
+MTP/gemma-4-12B-it-F16-MTP.gguf filter=lfs diff=lfs merge=lfs -text
+MTP/gemma-4-12B-it-Q4_0-MTP.gguf filter=lfs diff=lfs merge=lfs -text
+MTP/gemma-4-12B-it-Q8_0-MTP.gguf filter=lfs diff=lfs merge=lfs -text
+Tiểu[[:space:]]Ái[[:space:]]Test.mp4 filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/cublas64_12.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/cublasLt64_12.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/cudart64_12.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-base.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cpu-alderlake.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cpu-cannonlake.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cpu-cascadelake.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cpu-cooperlake.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cpu-haswell.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cpu-icelake.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cpu-ivybridge.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cpu-piledriver.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cpu-sandybridge.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cpu-sapphirerapids.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cpu-skylakex.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cpu-sse42.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cpu-x64.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cpu-zen4.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-cuda.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/ggml-rpc.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/libomp140.x86_64.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/llama-bench-impl.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/llama-cli-impl.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/llama-common.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/llama-completion-impl.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/llama-imatrix.exe filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/llama-perplexity-impl.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/llama-quantize-impl.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/llama-server-impl.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/llama-template-analysis.exe filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/llama-tts.exe filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/llama.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/mtmd.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama-cuda/rpc-server.exe filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-base.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-cpu-alderlake.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-cpu-cannonlake.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-cpu-cascadelake.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-cpu-cooperlake.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-cpu-haswell.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-cpu-icelake.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-cpu-ivybridge.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-cpu-piledriver.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-cpu-sandybridge.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-cpu-sapphirerapids.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-cpu-skylakex.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-cpu-sse42.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-cpu-x64.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-cpu-zen4.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/ggml-rpc.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/libomp140.x86_64.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/llama-bench-impl.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/llama-cli-impl.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/llama-common.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/llama-completion-impl.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/llama-imatrix.exe filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/llama-perplexity-impl.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/llama-quantize-impl.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/llama-server-impl.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/llama-template-analysis.exe filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/llama-tts.exe filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/llama.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/mtmd.dll filter=lfs diff=lfs merge=lfs -text
+tools/llama.cpp/rpc-server.exe filter=lfs diff=lfs merge=lfs -text
diff --git a/.gitignore b/.gitignore
new file mode 100644
index 0000000000000000000000000000000000000000..b0c9d2e2d084ed4cd49a130cd6b9f6eb9377d65b
--- /dev/null
+++ b/.gitignore
@@ -0,0 +1,6 @@
+# Local-only Hugging Face token (never commit/upload)
+hf_token.local
+hf_token.txt
+
+# GGUF weights (tải vào models/ — xem huggingface/scripts/download_models.py)
+models/*.gguf
diff --git a/.pytest_cache/.gitignore b/.pytest_cache/.gitignore
new file mode 100644
index 0000000000000000000000000000000000000000..bc1a1f6167d09c909aad37280b760bb715d0f1da
--- /dev/null
+++ b/.pytest_cache/.gitignore
@@ -0,0 +1,2 @@
+# Created by pytest automatically.
+*
diff --git a/.pytest_cache/CACHEDIR.TAG b/.pytest_cache/CACHEDIR.TAG
new file mode 100644
index 0000000000000000000000000000000000000000..fce15ad7eaa74e5682b644c84efb75334c112f95
--- /dev/null
+++ b/.pytest_cache/CACHEDIR.TAG
@@ -0,0 +1,4 @@
+Signature: 8a477f597d28d172789f06886806bc55
+# This file is a cache directory tag created by pytest.
+# For information about cache directory tags, see:
+# https://bford.info/cachedir/spec.html
diff --git a/.pytest_cache/README.md b/.pytest_cache/README.md
new file mode 100644
index 0000000000000000000000000000000000000000..b89018ced91c0a8af7f3f23ce8901870da89f3a0
--- /dev/null
+++ b/.pytest_cache/README.md
@@ -0,0 +1,8 @@
+# pytest cache directory #
+
+This directory contains data from the pytest's cache plugin,
+which provides the `--lf` and `--ff` options, as well as the `cache` fixture.
+
+**Do not** commit this to version control.
+
+See [the docs](https://docs.pytest.org/en/stable/how-to/cache.html) for more information.
diff --git a/.pytest_cache/v/cache/nodeids b/.pytest_cache/v/cache/nodeids
new file mode 100644
index 0000000000000000000000000000000000000000..7b5c525ff5efb88a8f882580c6fa9c116f1f3fb7
--- /dev/null
+++ b/.pytest_cache/v/cache/nodeids
@@ -0,0 +1,19 @@
+[
+ "tests/test_diarize_gender.py::test_age_group_empty_when_gender_unknown",
+ "tests/test_diarize_gender.py::test_age_group_unreliable_when_high_std",
+ "tests/test_diarize_gender.py::test_age_group_unreliable_when_too_few_windows",
+ "tests/test_diarize_gender.py::test_age_group_when_stable",
+ "tests/test_diarize_gender.py::test_reconcile_agreement_returns_gender",
+ "tests/test_diarize_gender.py::test_reconcile_ambiguous_pitch_needs_higher_conf",
+ "tests/test_diarize_gender.py::test_reconcile_child_passthrough",
+ "tests/test_diarize_gender.py::test_reconcile_f0_contradicts_female",
+ "tests/test_diarize_gender.py::test_reconcile_low_model_conf_returns_unknown",
+ "tests/test_diarize_gender.py::test_reconcile_model_pitch_conflict_returns_unknown",
+ "tests/test_relationship_lock.py::test_dominant_pair_picks_most_dialogue",
+ "tests/test_relationship_lock.py::test_ff_direct_anh_em_corrected_to_chi",
+ "tests/test_relationship_lock.py::test_kinship_source_beats_romance_lock",
+ "tests/test_relationship_lock.py::test_no_anh_em_lock_for_two_females",
+ "tests/test_relationship_lock.py::test_romance_lock_not_overwritten_by_translation_mom_child",
+ "tests/test_relationship_lock.py::test_same_age_locks_anh_em_despite_middle_aged_female",
+ "tests/test_relationship_lock.py::test_vision_lock_rejected_for_same_gender_pair"
+]
\ No newline at end of file
diff --git a/COLAB.md b/COLAB.md
new file mode 100644
index 0000000000000000000000000000000000000000..49a9dd68a3146f9d8fc722b5a02646cd0c6dc18b
--- /dev/null
+++ b/COLAB.md
@@ -0,0 +1,56 @@
+# Colab + Hugging Face
+
+Hướng dẫn chạy trên **Google Colab (L4)**.
+
+## Repo HF
+
+**https://huggingface.co/STBack23/gemma-srt-translate**
+
+## Cấu trúc Google Drive
+
+```
+Drive của tôi/
+└── Gemma/ ← folder chính
+ ├── Cache/ ← model GGUF + llama-server (tự cache)
+ └── Phim/ ← đặt video + SRT vào đây
+ ├── lam-chanh-anh.mp4
+ ├── lam-chanh-anh.srt
+ └── lam-chanh-anh.vi.srt ← file ra (sau khi dịch)
+```
+
+## Mở Colab
+
+| Cách | Link |
+|------|------|
+| **Open in Colab** | [colab/GemmaSRT_Colab.ipynb](https://huggingface.co/STBack23/gemma-srt-translate/blob/main/colab/GemmaSRT_Colab.ipynb) |
+| **Shortcut `/colab`** | [huggingface.co/.../colab](https://huggingface.co/STBack23/gemma-srt-translate/colab) |
+| **Bản mới nhất (nếu HF chưa kịp cập nhật)** | Colab → File → Upload notebook → `colab/GemmaSRT_Colab.ipynb` trên máy |
+
+> HF repo đôi khi chậm hơn file local. Nếu vẫn thấy `lam-chanh-anh` / `JOBS` cũ → upload file local hoặc chạy `hf auth login` rồi push lại (xem `huggingface/UPLOAD.md`).
+
+## Chạy
+
+1. Upload phim + SRT vào **`Gemma/Phim`**
+2. Sửa **`JOBS`** — chỉ cần tên file (không đuôi):
+
+```python
+JOBS = [
+ job("lam-chanh-anh"),
+ job("phim2"),
+]
+```
+
+3. Runtime → GPU **L4** → **Run all**
+
+### Tự ngắt sau khi xong (1 hoặc nhiều phim)
+
+```python
+DOWNLOAD_RESULTS = False # file đã trên Drive → không cần tải về
+AUTO_UNMOUNT_DRIVE = True
+AUTO_DISCONNECT_RUNTIME = True
+```
+
+## Local + Colab song song
+
+- **Local**: `GemmaSRT.bat` → phim A
+- **Colab**: notebook → phim trong `Gemma/Phim`
diff --git a/GemmaSRT.bat b/GemmaSRT.bat
new file mode 100644
index 0000000000000000000000000000000000000000..a36c29669d23781ca8392d5925e697caae3de09e
--- /dev/null
+++ b/GemmaSRT.bat
@@ -0,0 +1,5 @@
+@echo off
+rem Khoi chay Gemma SRT Translate (khong hien console)
+set "ROOT=%~dp0"
+cd /d "%ROOT%"
+start "" pythonw "%ROOT%tools\gemma_srt_app.py"
diff --git a/GemmaSRT.vbs b/GemmaSRT.vbs
new file mode 100644
index 0000000000000000000000000000000000000000..8da5ec69a573476daa71109af875f18a2b719967
--- /dev/null
+++ b/GemmaSRT.vbs
@@ -0,0 +1,4 @@
+' Chay GemmaSRT.bat o che do an (khong nhay cua so console)
+Set sh = CreateObject("WScript.Shell")
+dir = Left(WScript.ScriptFullName, InStrRev(WScript.ScriptFullName, "\"))
+sh.Run """" & dir & "GemmaSRT.bat""", 0, False
diff --git a/MTP/README.md b/MTP/README.md
new file mode 100644
index 0000000000000000000000000000000000000000..1852abcaad46b860cfd241def3ff5d45a94495d8
--- /dev/null
+++ b/MTP/README.md
@@ -0,0 +1,58 @@
+# Gemma 4 12B QAT MTP drafter
+
+Multi-Token Prediction (MTP) drafter for `unsloth/gemma-4-12B-it-qat-GGUF`. It runs as a speculative draft model that shares the target's KV cache and speeds up text generation. The drafter is the Gemma 4 12B QAT assistant and pairs with the QAT model unchanged.
+
+Verified on a single B200 against the `gemma-4-12B-it-qat-UD-Q4_K_XL.gguf` target with `-hf` auto-discovery: draft acceptance 0.51.
+
+MTP was merged into llama.cpp on 2026-06-07 (PR ggml-org/llama.cpp#23398). You need a llama.cpp build **>= b9549** (`llama-server --version`). Older builds cannot load these (arch `gemma4-assistant`).
+
+## Files
+
+The recommended drafter is a **smart Q4_0**: the native 4-bit QAT drafter (about 97% of its weights are byte-exact on the int4 grid), near-lossless versus higher precision while roughly half the size. It sits at the repo root as `mtp-gemma-4-12B-it.gguf` so `-hf` finds it automatically, and the same file plus higher-precision drafters are in `MTP/`:
+
+- `mtp-gemma-4-12B-it.gguf` (repo root, smart Q4_0, recommended; used by `-hf`)
+- `MTP/gemma-4-12B-it-Q4_0-MTP.gguf` (same smart Q4_0)
+- `MTP/gemma-4-12B-it-Q8_0-MTP.gguf`
+- `MTP/gemma-4-12B-it-BF16-MTP.gguf`
+- `MTP/gemma-4-12B-it-F16-MTP.gguf`
+
+## Build llama.cpp
+
+```bash
+git clone https://github.com/ggml-org/llama.cpp
+cd llama.cpp
+
+# CUDA build. Set the arch for your GPU: 89 (RTX 4090), 90 (H100), 100 (B200).
+cmake -B build -DGGML_CUDA=ON -DCMAKE_CUDA_ARCHITECTURES=90
+cmake --build build --config Release -j --target llama-server
+```
+
+## Run, the easy way
+
+A recent llama.cpp finds the drafter automatically from the root `mtp-` file, so `-hf` is all you need. No `--model-draft`.
+
+```bash
+./build/bin/llama-server \
+ -hf unsloth/gemma-4-12B-it-qat-GGUF:UD-Q4_K_XL \
+ --spec-type draft-mtp --spec-draft-n-max 4 \
+ -ngl 999 -fa on
+```
+
+If your build is too old to auto-discover the sibling, use the explicit form below.
+
+## Run with an explicit drafter
+
+Use this to choose a precision or point at a local file.
+
+```bash
+hf download unsloth/gemma-4-12B-it-qat-GGUF gemma-4-12B-it-qat-UD-Q4_K_XL.gguf --local-dir .
+hf download unsloth/gemma-4-12B-it-qat-GGUF MTP/gemma-4-12B-it-Q8_0-MTP.gguf --local-dir .
+
+./build/bin/llama-server \
+ -m gemma-4-12B-it-qat-UD-Q4_K_XL.gguf \
+ --model-draft MTP/gemma-4-12B-it-Q8_0-MTP.gguf \
+ --spec-type draft-mtp --spec-draft-n-max 4 \
+ -ngl 999 -fa on
+```
+
+Multi GPU: add `--spec-draft-device CUDA0 -sm layer`. The drafter pairs with any quant of the 12B QAT model. Quantized KV cache works.
diff --git a/MTP/gemma-4-12B-it-BF16-MTP.gguf b/MTP/gemma-4-12B-it-BF16-MTP.gguf
new file mode 100644
index 0000000000000000000000000000000000000000..5907997d6ef698161bab7de2002487504f3d3b41
--- /dev/null
+++ b/MTP/gemma-4-12B-it-BF16-MTP.gguf
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:320b1dfe0dc0af8380c1a74b4992177841a940da92015b9c672e2590b20c17de
+size 861537344
diff --git a/MTP/gemma-4-12B-it-F16-MTP.gguf b/MTP/gemma-4-12B-it-F16-MTP.gguf
new file mode 100644
index 0000000000000000000000000000000000000000..a93593fb5bbc7d115e7639c069f4e40b0244602b
--- /dev/null
+++ b/MTP/gemma-4-12B-it-F16-MTP.gguf
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:bb7a268dc6518c37b71095601140ef930630f16fa3b9be91311a625013235601
+size 861537344
diff --git a/MTP/gemma-4-12B-it-Q4_0-MTP.gguf b/MTP/gemma-4-12B-it-Q4_0-MTP.gguf
new file mode 100644
index 0000000000000000000000000000000000000000..92a8302495be3377e4b7ae967324c4eafff661aa
--- /dev/null
+++ b/MTP/gemma-4-12B-it-Q4_0-MTP.gguf
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:c50c91c35f04903815b2e8930cbb8c8c5bee0e1aa00748c30a7b8ff05d2310b4
+size 253707328
diff --git a/MTP/gemma-4-12B-it-Q8_0-MTP.gguf b/MTP/gemma-4-12B-it-Q8_0-MTP.gguf
new file mode 100644
index 0000000000000000000000000000000000000000..70e31011ac0ca97cf3f1ca3ce74c72a98e365f14
--- /dev/null
+++ b/MTP/gemma-4-12B-it-Q8_0-MTP.gguf
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:3c9460bc9f6143f232c9447bc772bf6f6db974b57129298381cf0cb702bd586b
+size 465126464
diff --git a/README.md b/README.md
index bcdab3d85d5123ce8157b684fae3c63d1ebb60b9..74c156fd8346684cd68267e0b4b1d97fbf256b34 100644
--- a/README.md
+++ b/README.md
@@ -1,139 +1,578 @@
---
+library_name: transformers
license: apache-2.0
-language:
- - vi
- - en
+license_link: https://ai.google.dev/gemma/docs/gemma_4_license
+pipeline_tag: any-to-any
+base_model: google/gemma-4-12B-it-qat-q4_0-unquantized
tags:
- - subtitles
- - srt
- - translation
- - gemma4
- - vision
- - llama.cpp
-pipeline_tag: translation
-library_name: gemma-srt-translate
+- gemma4
+- unsloth
+- gemma
+- google
---
+# Read our How to [Run Gemma 4 QAT Guide!](https://unsloth.ai/docs/models/gemma-4/qat)
+
+
+
+
+
+
+## Run with MTP (speculative decoding)
+
+This model ships a Multi-Token Prediction drafter at the repo root (`mtp-gemma-4-12B-it.gguf`, a near-lossless smart Q4_0). A recent llama.cpp auto-discovers it from `-hf`, so you do not pass `--model-draft`:
-# Gemma SRT Translate
+```bash
+./build/bin/llama-server \
+ -hf unsloth/gemma-4-12B-it-qat-GGUF:UD-Q4_K_XL \
+ --spec-type draft-mtp --spec-draft-n-max 4 \
+ -ngl 999 -fa on
+```
+
+The drafter shares the target's KV cache and does not change the output (the target verifies every drafted token). See the `MTP/` folder for the other precisions and explicit usage.
+
+
+
+

+
+
+
+
+ Hugging Face |
+ GitHub |
+ Launch Blog |
+ Documentation
+
+ License: Apache 2.0 | Authors: Google DeepMind
+
+
+> [!Note]
+> This model card is for the new versions of the Gemma 4 family optimized with Quantization-Aware Training (QAT), which allows preserving similar quality to bfloat16 while dramatically reducing the memory requirements to load the model.
+> Four versions of the QAT checkpoints are available:
+> * **Unquantized QAT checkpoints** (Q4_0): Half-precision weights extracted from the QAT pipeline, ideal for custom downstream compilation and research. Available for Gemma 4 E2B, E4B, 12B, 26B A4B, and 31B, and their drafter models.
+> * **GGUF** (Q4_0): Ready-to-deploy formats for broad ecosystem compatibility. Available for Gemma 4 E2B, E4B, 12B, 26B A4B, and 31B.
+> * **Mobile-optimized** (wNa8o8): A custom schema engineered explicitly for mobile hardware efficiency. It features targeted 2-bit decoding layers, optimized KV caches, and static activations to maximize VRAM savings. Available for Gemma 4 E2B and E4B.
+> * **Compressed Tensors** (w4a16): QAT checkpoints serialized in the compressed-tensors format for native, optimized inference with vLLM. Available for Gemma 4 E2B, E4B, 12B, and 31B.
+
+Gemma is a family of open models built by Google DeepMind. Gemma 4 models are multimodal, handling text and image input (with audio supported on E2B, E4B, and 12B) and generating text output. This release includes open-weights models in both pre-trained and instruction-tuned variants. Gemma 4 features a context window of up to 256K tokens and maintains multilingual support in over 140 languages.
+
+Featuring both Dense and Mixture-of-Experts (MoE) architectures, Gemma 4 is well-suited for tasks like text generation, coding, and reasoning. The models are available in five distinct sizes: **E2B**, **E4B**, **12B**, **26B A4B**, and **31B**. Their diverse sizes make them deployable in environments ranging from high-end phones to laptops and servers, democratizing access to state-of-the-art AI.
+
+Gemma 4 introduces key **capability and architectural advancements**:
+
+* **Reasoning** – All models in the family are designed as highly capable reasoners, with configurable thinking modes.
+
+* **Extended Multimodalities** – Processes Text, Image with variable aspect ratio and resolution support (all models), Video, and Audio (featured natively on the E2B, E4B, and 12B models).
+
+* **Diverse & Efficient Architectures** – Offers Dense and Mixture-of-Experts (MoE) variants of different sizes for scalable deployment.
+
+* **Optimized for On-Device** – Smaller models are specifically designed for efficient local execution on laptops and mobile devices.
+
+* **Increased Context Window** – The small models feature a 128K context window, while the medium models support 256K.
+
+* **Enhanced Coding & Agentic Capabilities** – Achieves notable improvements in coding benchmarks alongside native function-calling support, powering highly capable autonomous agents.
+
+* **Native System Prompt Support** – Gemma 4 introduces native support for the `system` role, enabling more structured and controllable conversations.
+
+## **Models Overview**
+
+Gemma 4 models are designed to deliver frontier-level performance at each size, targeting deployment scenarios from mobile and edge devices (E2B, E4B) to consumer GPUs and workstations (12B, 26B A4B, 31B). They are well-suited for reasoning, agentic workflows, coding, and multimodal understanding.
+
+The models employ a hybrid attention mechanism that interleaves local sliding window attention with full global attention, ensuring the final layer is always global. This hybrid design delivers the processing speed and low memory footprint of a lightweight model without sacrificing the deep awareness required for complex, long-context tasks. To optimize memory for long contexts, global layers feature unified Keys and Values, and apply Proportional RoPE (p-RoPE).
+
+### Dense Models
+
+| Property | E2B | E4B | 12B Unified | 31B Dense |
+| :---- | :---- | :---- | :---- | :---- |
+| **Total Parameters** | 2.3B effective
(5.1B with embeddings) | 4.5B effective
(8B with embeddings) | 11.95B | 30.7B |
+| **Layers** | 35 | 42 | 48 | 60 |
+| **Sliding Window** | 512 tokens | 512 tokens | 1024 tokens | 1024 tokens |
+| **Context Length** | 128K tokens | 128K tokens | 256K tokens | 256K tokens |
+| **Vocabulary Size** | 262K | 262K | 262K | 262K |
+| **Supported Modalities** | Text, Image, Audio | Text, Image, Audio | Text, Image, Audio | Text, Image |
+| **Vision Encoder Parameters** | *~150M* | *~150M* | - | *~550M* |
+| **Audio Encoder Parameters** | *~300M* | *~300M* | - | No Audio |
+
+The "E" in E2B and E4B stands for "effective" parameters. The smaller models incorporate Per-Layer Embeddings (PLE) to maximize parameter efficiency in on-device deployments. Rather than adding more layers or parameters to the model, PLE gives each decoder layer its own small embedding for every token. These embedding tables are large but are only used for quick lookups, which is why the effective parameter count is much smaller than the total.
+
+The "Unified" in Gemma 4 12B Unified refers to its encoder-free architecture. Other Gemma 4 models use dedicated encoders to process multimodal data before passing it to the LLM. Gemma 4 12B eliminates these encoders entirely, projecting raw image patches and audio waveforms directly into the LLM's embedding space through lightweight linear layers. This unified approach means all modalities flow straight into a single decoder-only transformer, reducing multimodal latency and allowing the entire model to be fine-tuned in one pass.
+
+### Mixture-of-Experts (MoE) Model
+
+| Property | 26B A4B MoE |
+| :---- | :---- |
+| **Total Parameters** | 25.2B |
+| **Active Parameters** | 3.8B |
+| **Layers** | 30 |
+| **Sliding Window** | 1024 tokens |
+| **Context Length** | 256K tokens |
+| **Vocabulary Size** | 262K |
+| **Expert Count** | 8 active / 128 total and 1 shared |
+| **Supported Modalities** | Text, Image |
+| **Vision Encoder Parameters** | *~550M* |
-Dịch và sửa phụ đề SRT bằng **Gemma 4 12B vision** (llama-server + MTP).
-Repo này chứa **script + notebook Colab**; model GGUF tải từ [`unsloth/gemma-4-12B-it-qat-GGUF`](https://huggingface.co/unsloth/gemma-4-12B-it-qat-GGUF).
+The "A" in 26B A4B stands for "active parameters" in contrast to the total number of parameters the model contains. By only activating a 4B subset of parameters during inference, the Mixture-of-Experts model runs much faster than its 26B total might suggest. This makes it an excellent choice for fast inference compared to the dense 31B model since it runs almost as fast as a 4B-parameter model.
-## Cấu trúc repo
+## **Benchmark Results**
+These models were evaluated against a large collection of different datasets and metrics to cover different aspects of text generation. Evaluation results marked in the table are for instruction-tuned models.
+
+| | Gemma 4 31B | Gemma 4 26B A4B | Gemma 4 12B Unified | Gemma 4 E4B | Gemma 4 E2B | Gemma 3 27B (no think) |
+| :---- | :---- | :---- | :---- | :---- | :---- | :---- |
+| MMLU Pro | 85.2% | 82.6% | 77.2% | 69.4% | 60.0% | 67.6% |
+| AIME 2026 no tools | 89.2% | 88.3% | 77.5% | 42.5% | 37.5% | 20.8% |
+| LiveCodeBench v6 | 80.0% | 77.1% | 72.0% | 52.0% | 44.0% | 29.1% |
+| Codeforces ELO | 2150 | 1718 | 1659 | 940 | 633 | 110 |
+| GPQA Diamond | 84.3% | 82.3% | 78.8% | 58.6% | 43.4% | 42.4% |
+| Tau2 (average over 3) | 76.9% | 68.2% | 69.0% | 42.2% | 24.5% | 16.2% |
+| HLE no tools | 19.5% | 8.7% | 5.2% | - | - | - |
+| HLE with search | 26.5% | 17.2% | - | - | - | - |
+| BigBench Extra Hard | 74.4% | 64.8% | 53.0% | 33.1% | 21.9% | 19.3% |
+| MMMLU | 88.4% | 86.3% | 83.4% | 76.6% | 67.4% | 70.7% |
+| **Vision** | | | | | | |
+| MMMU Pro | 76.9% | 73.8% | 69.1% | 52.6% | 44.2% | 49.7% |
+| OmniDocBench 1.5 (average edit distance, lower is better) | 0.131 | 0.149 | 0.164 | 0.181 | 0.290 | 0.365 |
+| MATH-Vision | 85.6% | 82.4% | 79.7% | 59.5% | 52.4% | 46.0% |
+| MedXPertQA MM | 61.3% | 58.1% | 48.7% | 28.7% | 23.5% | - |
+| **Audio** | | | | | | |
+| CoVoST | - | - | 38.5* | 35.54 | 33.47 | - |
+| FLEURS (lower is better) | - | - | 0.069* | 0.08 | 0.09 | - |
+| **Long Context** | | | | | | |
+| MRCR v2 8 needle 128k (average) | 66.4% | 44.1% | 43.4% | 25.4% | 19.1% | 13.5% |
+
+*Excluding Chinese language.
+
+## **Core Capabilities**
+
+Gemma 4 models handle a broad range of tasks across text, vision, and audio. Key capabilities include:
+
+* **Thinking** – Built-in reasoning mode that lets the model think step-by-step before answering.
+* **Long Context** – Context windows of up to 128K tokens (E2B/E4B) and 256K tokens (12B, 26B A4B/31B).
+* **Image Understanding** – Object detection, Document/PDF parsing, screen and UI understanding, chart comprehension, OCR (including multilingual), handwriting recognition, and pointing. Images can be processed at variable aspect ratios and resolutions.
+* **Video Understanding** – Analyze video by processing sequences of frames.
+* **Interleaved Multimodal Input** – Freely mix text and images in any order within a single prompt.
+* **Function Calling** – Native support for structured tool use, enabling agentic workflows.
+* **Coding** – Code generation, completion, and correction.
+* **Multilingual** – Out-of-the-box support for 35+ languages, pre-trained on 140+ languages.
+* **Audio** (E2B, E4B, and 12B only) – Automatic speech recognition (ASR) and speech-to-translated-text translation across multiple languages.
+
+
+## Getting Started
+
+You can use all Gemma 4 models with the latest version of Transformers. To get started, install the necessary dependencies in your environment:
+
+`pip install -U transformers torch accelerate`
+
+Once you have everything installed, you can proceed to load the model with the code below:
+
+```python
+from transformers import AutoProcessor, AutoModelForMultimodalLM
+
+MODEL_ID = "google/gemma-4-12B-it"
+
+# Load model
+processor = AutoProcessor.from_pretrained(MODEL_ID)
+model = AutoModelForMultimodalLM.from_pretrained(
+ MODEL_ID,
+ dtype="auto",
+ device_map="auto"
+)
```
-YOUR_USERNAME/gemma-srt-translate/
-├── README.md ← file này
-├── config.yaml ← cấu hình mặc định
-├── translate_srt.py ← pipeline chính
-├── diarize_audio.py ← Pass 0: phân tích giọng nói (tùy chọn)
-├── requirements-diarize.txt ← dependency cho Pass 0
-├── colab/
-│ └── GemmaSRT_Colab.ipynb ← chạy trên Google Colab (L4 24GB)
-└── scripts/
- ├── download_models.py ← tải GGUF từ Unsloth
- └── build_llama_server.sh ← build llama-server Linux (Colab)
+
+Once the model is loaded, you can start generating output:
+
+```python
+# Prompt
+messages = [
+ {"role": "system", "content": "You are a helpful assistant."},
+ {"role": "user", "content": "Write a short joke about saving RAM."},
+]
+
+# Process input
+inputs = processor.apply_chat_template(
+ messages,
+ tokenize=True,
+ return_dict=True,
+ return_tensors="pt",
+ add_generation_prompt=True,
+ enable_thinking=False
+).to(model.device)
+input_len = inputs["input_ids"].shape[-1]
+
+# Generate output
+outputs = model.generate(**inputs, max_new_tokens=1024)
+response = processor.decode(outputs[0][input_len:], skip_special_tokens=False)
+
+# Parse output
+processor.parse_response(response)
```
-## Chạy trên Google Colab (khuyến nghị L4 24GB)
+To enable reasoning, set `enable_thinking=True` and the `parse_response` function will take care of parsing the thinking output.
-Repo **Private** → **không** mở được bằng nút Open in Colab trên HF (lỗi 401).
+Below, you will also find snippets for processing audio (E2B, E4B, 12B only), images, and video alongside text:
-**Cách mở:** tải [colab/GemmaSRT_Colab.ipynb](https://huggingface.co/STBack23/gemma-srt-translate/blob/main/colab/GemmaSRT_Colab.ipynb) → Colab **File → Upload notebook** → dán **`HF_TOKEN`** vào cell cấu hình → lưu copy lên Drive `Gemma/` để dùng lại.
+
+Code for processing Audio
-Repo này nằm trên **Hugging Face**, không phải GitHub — link `colab.research.google.com/github/...` sẽ lỗi 404.
+Make sure to install the following packages:
-**Cách mở (chọn một):**
+`pip install -U transformers torch torchvision librosa accelerate`
-1. **Nút Open in Colab trên HF** (khuyến nghị):
- [colab/GemmaSRT_Colab.ipynb](https://huggingface.co/STBack23/gemma-srt-translate/blob/main/colab/GemmaSRT_Colab.ipynb) → bấm **Open in Colab**
+You can then load the model with the code below:
-2. **Shortcut HF `/colab`**:
- [https://huggingface.co/STBack23/gemma-srt-translate/colab](https://huggingface.co/STBack23/gemma-srt-translate/colab)
+```python
+from transformers import AutoProcessor, AutoModelForMultimodalLM
-3. **Colab → File → Upload notebook** → tải file `.ipynb` từ HF về rồi upload
+MODEL_ID = "google/gemma-4-12B-it"
-Sau khi mở notebook:
+# Load model
+processor = AutoProcessor.from_pretrained(MODEL_ID)
+model = AutoModelForMultimodalLM.from_pretrained(
+ MODEL_ID,
+ dtype="auto",
+ device_map="auto"
+)
+```
-1. Runtime → **Change runtime type** → GPU **L4** (hoặc T4/A100)
-2. Chạy tuần tự các cell (lần đầu ~15–20 phút: build llama-server + tải model)
-3. Upload video + SRT hoặc trỏ vào Google Drive
-4. Tải file `*.vi.srt` về máy
+Once the model is loaded, you can start generating output by directly referencing the audio URL in the prompt:
+
+
+```python
+# Prompt - add audio after text
+messages = [
+ {
+ "role": "user",
+ "content": [
+ {"type": "text", "text": "Transcribe the following speech segment in its original language. Follow these specific instructions for formatting the answer:\n* Only output the transcription, with no newlines.\n* When transcribing numbers, write the digits, i.e. write 1.7 and not one point seven, and write 3 instead of three."},
+ {"type": "audio", "audio": "https://raw.githubusercontent.com/google-gemma/cookbook/refs/heads/main/apps/sample-data/journal1.wav"},
+ ]
+ }
+]
+
+# Process input
+inputs = processor.apply_chat_template(
+ messages,
+ tokenize=True,
+ return_dict=True,
+ return_tensors="pt",
+ add_generation_prompt=True,
+).to(model.device)
+input_len = inputs["input_ids"].shape[-1]
+
+# Generate output
+outputs = model.generate(**inputs, max_new_tokens=512)
+response = processor.decode(outputs[0][input_len:], skip_special_tokens=False)
+
+# Parse output
+processor.parse_response(response)
+```
-[](https://huggingface.co/STBack23/gemma-srt-translate/colab)
+
-## Chạy trên máy local (Windows)
+
+Code for processing Images
-Repo HF **không** chứa `llama-server.exe` — dùng bản Windows trong project gốc:
+Make sure to install the following packages:
-```powershell
-.\run-translate.ps1 -Video "phim.mp4" -InputSrt "phim.srt" -OutputSrt "phim.vi.srt"
-```
-Hoặc mở GUI: `GemmaSRT.bat`
+`pip install -U transformers torch torchvision accelerate`
-## Tải model thủ công
+You can then load the model with the code below:
-```bash
-pip install huggingface_hub
-python scripts/download_models.py --dest ./models
+```python
+from transformers import AutoProcessor, AutoModelForMultimodalLM
+
+MODEL_ID = "google/gemma-4-12B-it"
+
+# Load model
+processor = AutoProcessor.from_pretrained(MODEL_ID)
+model = AutoModelForMultimodalLM.from_pretrained(
+ MODEL_ID,
+ dtype="auto",
+ device_map="auto"
+)
```
-File tải về (~12–15 GB):
+Once the model is loaded, you can start generating output by directly referencing the image URL in the prompt:
+
+
+```python
+# Prompt - add image before text
+messages = [
+ {
+ "role": "user", "content": [
+ {"type": "image", "url": "https://raw.githubusercontent.com/google-gemma/cookbook/refs/heads/main/apps/sample-data/GoldenGate.png"},
+ {"type": "text", "text": "What is shown in this image?"}
+ ]
+ }
+]
+
+# Process input
+inputs = processor.apply_chat_template(
+ messages,
+ tokenize=True,
+ return_dict=True,
+ return_tensors="pt",
+ add_generation_prompt=True,
+).to(model.device)
+input_len = inputs["input_ids"].shape[-1]
+
+# Generate output
+outputs = model.generate(**inputs, max_new_tokens=512)
+response = processor.decode(outputs[0][input_len:], skip_special_tokens=False)
+
+# Parse output
+processor.parse_response(response)
+```
-| File | Mục đích |
-|------|----------|
-| `gemma-4-12B-it-qat-UD-Q4_K_XL.gguf` | Model chính |
-| `mmproj-F16.gguf` | Vision projector |
-| `mtp-gemma-4-12B-it.gguf` | MTP draft (tăng tốc) |
+
-## Pipeline
-1. Parse SRT → gom scene
-2. **Pass 0** *(tùy chọn)*: phân tích giọng nói (pyannote) → "ai nói câu nào" + giới tính/độ tuổi → gán vào từng cue để xưng hô nhất quán
-3. **Pass 1**: cắt frame (ffmpeg) → OCR/sửa SRT gốc (vision)
-4. **Pass 2**: dịch theo ngữ cảnh + xưng hô
-5. Ghi SRT dịch, SRT đã sửa, báo cáo JSON
+
+Code for processing Videos
-### Pass 0 — phân tích giọng nói (diarization)
+Make sure to install the following packages:
-Bật bằng `--diarize` (CLI) hoặc `ENABLE_DIARIZE = True` (Colab). Cần thêm:
+`pip install -U transformers torch torchvision librosa accelerate`
-```bash
-pip install -r requirements-diarize.txt
+You can then load the model with the code below:
+
+```python
+from transformers import AutoProcessor, AutoModelForMultimodalLM
+
+MODEL_ID = "google/gemma-4-12B-it"
+
+# Load model
+processor = AutoProcessor.from_pretrained(MODEL_ID)
+model = AutoModelForMultimodalLM.from_pretrained(
+ MODEL_ID,
+ dtype="auto",
+ device_map="auto"
+)
```
-Token HF phải bấm **Agree** điều kiện các model gated:
-[community-1](https://hf.co/pyannote/speaker-diarization-community-1) ·
-[3.1](https://hf.co/pyannote/speaker-diarization-3.1) ·
-[segmentation-3.0](https://hf.co/pyannote/segmentation-3.0).
-`--gender-method model` dùng [audeering wav2vec2 age/gender](https://hf.co/audeering/wav2vec2-large-robust-24-ft-age-gender) (~1GB).
+Once the model is loaded, you can start generating output by directly referencing the video URL in the prompt:
+
+
+```python
+# Prompt - add video before text
+messages = [
+ {
+ 'role': 'user',
+ 'content': [
+ {"type": "video", "video": "https://github.com/bebechien/gemma/raw/refs/heads/main/videos/ForBiggerBlazes.mp4"},
+ {'type': 'text', 'text': 'Describe this video.'}
+ ]
+ }
+]
+
+# Process input
+inputs = processor.apply_chat_template(
+ messages,
+ tokenize=True,
+ return_dict=True,
+ return_tensors="pt",
+ add_generation_prompt=True,
+).to(model.device)
+input_len = inputs["input_ids"].shape[-1]
+
+# Generate output
+outputs = model.generate(**inputs, max_new_tokens=512)
+response = processor.decode(outputs[0][input_len:], skip_special_tokens=False)
+
+# Parse output
+processor.parse_response(response)
+```
-## VRAM
+
-| GPU | 1 phim | 2 phim song song |
-|-----|--------|------------------|
-| RTX 4060 Ti 16GB | ✅ | ❌ |
-| Colab L4 24GB | ✅ | ❌ (dùng local + Colab = 2 phim độc lập) |
-## Cấu hình
-Chỉnh `config.yaml` hoặc tham số CLI:
+## **Best Practices**
-```bash
-python translate_srt.py \
- --video phim.mp4 \
- --input-srt phim.srt \
- --output-srt phim.vi.srt \
- --target-lang Vietnamese \
- --skip-correction # bỏ pass OCR (nhanh hơn)
+For the best performance, use these configurations and best practices:
+
+### 1. Sampling Parameters
+
+Use the following standardized sampling configuration across all use cases:
+
+* `temperature=1.0`
+* `top_p=0.95`
+* `top_k=64`
+
+### 2. Thinking Mode Configuration
+
+Compared to Gemma 3, the models use standard `system`, `assistant`, and `user` roles. To properly manage the thinking process, use the following control tokens:
+
+* **Trigger Thinking:** Thinking is enabled by including the `<|think|>` token at the start of the system prompt. To disable thinking, remove the token.
+* **Standard Generation:** When thinking is enabled, the model will output its internal reasoning followed by the final answer using this structure:
+ `<|channel>thought\n`**[Internal reasoning]**``
+* **Disabled Thinking Behavior:** For all models except for the E2B and E4B variants, if thinking is disabled, the model will still generate the tags but with an empty thought block:
+ `<|channel>thought\n`**[Final answer]**
+
+> [!Note]
+> Note that many libraries like Transformers and llama.cpp handle the complexities of the chat template for you.
+
+### 3. Multi-Turn Conversations
+
+* **No Thinking Content in History**: In multi-turn conversations, the historical model output should only include the final response. Thoughts from previous model turns must *not be added* before the next user turn begins.
+
+### 4. Modality order
+
+For optimal performance with multimodal inputs, place:
+
+* Image content **before** the text in your prompt.
+* Audio content **after** the text in your prompt.
+
+### 5. Variable Image Resolution
+
+Aside from variable aspect ratios, Gemma 4 supports variable image resolution through a configurable visual token budget, which controls how many tokens are used to represent an image. A higher token budget preserves more visual detail at the cost of additional compute, while a lower budget enables faster inference for tasks that don't require fine-grained understanding.
+
+* The supported token budgets are: **70**, **140**, **280**, **560**, and **1120**.
+ * Use *lower budgets* for classification, captioning, or video understanding, where faster inference and processing many frames outweigh fine-grained detail.
+ * Use *higher budgets* for tasks like OCR, document parsing, or reading small text.
+
+### 6. Audio
+
+Use the following prompt structures for audio processing:
+
+* **Audio Speech Recognition (ASR)**
+
+```text
+Transcribe the following speech segment in {LANGUAGE} into {LANGUAGE} text.
+
+Follow these specific instructions for formatting the answer:
+* Only output the transcription, with no newlines.
+* When transcribing numbers, write the digits, i.e. write 1.7 and not one point seven, and write 3 instead of three.
+```
+
+* **Automatic Speech Translation (AST)**
+
+```text
+Transcribe the following speech segment in {SOURCE_LANGUAGE}, then translate it into {TARGET_LANGUAGE}.
+When formatting the answer, first output the transcription in {SOURCE_LANGUAGE}, then one newline, then output the string '{TARGET_LANGUAGE}: ', then the translation in {TARGET_LANGUAGE}.
```
-## Model gốc
+### 7. Audio and Video Length
+
+All models support image inputs and can process videos as frames whereas the E2B, E4B, and 12B models also support audio inputs. Audio supports a maximum length of 30 seconds. Video supports a maximum of 60 seconds assuming the images are processed at one frame per second.
+
+## **Model Data**
+
+Data used for model training and how the data was processed.
+
+### **Training Dataset**
+
+Our pre-training dataset is a large-scale, diverse collection of data encompassing a wide range of domains and modalities, which includes web documents, code, images, audio, with a cutoff date of January 2025. Here are the key components:
+
+* **Web Documents**: A diverse collection of web text ensures the model is exposed to a broad range of linguistic styles, topics, and vocabulary. The training dataset includes content in over 140 languages.
+* **Code**: Exposing the model to code helps it to learn the syntax and patterns of programming languages, which improves its ability to generate code and understand code-related questions.
+* **Mathematics**: Training on mathematical text helps the model learn logical reasoning, symbolic representation, and to address mathematical queries.
+* **Images**: A wide range of images enables the model to perform image analysis and visual data extraction tasks.
+
+The combination of these diverse data sources is crucial for training a powerful multimodal model that can handle a wide variety of different tasks and data formats.
+
+### **Data Preprocessing**
+
+Here are the key data cleaning and filtering methods applied to the training data:
+
+* **CSAM Filtering**: Rigorous CSAM (Child Sexual Abuse Material) filtering was applied at multiple stages in the data preparation process to ensure the exclusion of harmful and illegal content.
+* **Sensitive Data Filtering**: As part of making Gemma pre-trained models safe and reliable, automated techniques were used to filter out certain personal information and other sensitive data from training sets.
+* **Additional methods**: Filtering based on content quality and safety in line with [our policies](https://ai.google/static/documents/ai-responsibility-update-published-february-2025.pdf).
+
+## **Ethics and Safety**
+
+As open models become central to enterprise infrastructure, provenance and security are paramount. Developed by Google DeepMind, Gemma 4 undergoes the same rigorous safety evaluations as our proprietary Gemini models.
+
+### **Evaluation Approach**
+
+Gemma 4 models were developed in partnership with internal safety and responsible AI teams. A range of automated as well as human evaluations were conducted to help improve model safety. These evaluations align with [Google’s AI principles](https://ai.google/principles/), as well as safety policies, which aim to prevent our generative AI models from generating harmful content, including:
+
+* Content related to child sexual abuse material and exploitation
+* Dangerous content (e.g., promoting suicide, or instructing in activities that could cause real-world harm)
+* Sexually explicit content
+* Hate speech (e.g., dehumanizing members of protected groups)
+* Harassment (e.g., encouraging violence against people)
+
+### **Evaluation Results**
+
+For all areas of safety testing, we saw major improvements in all categories of content safety relative to previous Gemma models. Overall, Gemma 4 models significantly outperform Gemma 3 and 3n models in improving safety, while keeping unjustified refusals low. All testing was conducted without safety filters to evaluate the model capabilities and behaviors. For both text-to-text and image-to-text, and across all model sizes, the model produced minimal policy violations, and showed significant improvements over previous Gemma models' performance.
+
+## **Usage and Limitations**
+
+These models have certain limitations that users should be aware of.
+
+### **Intended Usage**
+
+Multimodal models (capable of processing vision, language, and/or audio) have a wide range of applications across various industries and domains. The following list of potential uses is not comprehensive. The purpose of this list is to provide contextual information about the possible use-cases that the model creators considered as part of model training and development.
+
+* **Content Creation and Communication**
+ * **Text Generation**: These models can be used to generate creative text formats such as poems, scripts, code, marketing copy, and email drafts.
+ * **Chatbots and Conversational AI**: Power conversational interfaces for customer service, virtual assistants, or interactive applications.
+ * **Text Summarization**: Generate concise summaries of a text corpus, research papers, or reports.
+ * **Image Data Extraction**: These models can be used to extract, interpret, and summarize visual data for text communications.
+ * **Audio Processing and Interaction**: The E2B, E4B, and 12B models can analyze and interpret audio inputs, enabling voice-driven interactions and transcriptions.
+* **Research and Education**
+ * **Natural Language Processing (NLP) and VLM Research**: These models can serve as a foundation for researchers to experiment with VLM and NLP techniques, develop algorithms, and contribute to the advancement of the field.
+ * **Language Learning Tools**: Support interactive language learning experiences, aiding in grammar correction or providing writing practice.
+ * **Knowledge Exploration**: Assist researchers in exploring large bodies of text by generating summaries or answering questions about specific topics.
+
+### **Limitations**
+
+* **Training Data**
+ * The quality and diversity of the training data significantly influence the model's capabilities. Biases or gaps in the training data can lead to limitations in the model's responses.
+ * The scope of the training dataset determines the subject areas the model can handle effectively.
+* **Context and Task Complexity**
+ * Models perform well on tasks that can be framed with clear prompts and instructions. Open-ended or highly complex tasks might be challenging.
+ * A model's performance can be influenced by the amount of context provided (longer context generally leads to better outputs, up to a certain point).
+* **Language Ambiguity and Nuance**
+ * Natural language is inherently complex. Models might struggle to grasp subtle nuances, sarcasm, or figurative language.
+* **Factual Accuracy**
+ * Models generate responses based on information they learned from their training datasets, but they are not knowledge bases. They may generate incorrect or outdated factual statements.
+* **Common Sense**
+ * Models rely on statistical patterns in language. They might lack the ability to apply common sense reasoning in certain situations.
+
+### **Ethical Considerations and Risks**
+
+The development of vision-language models (VLMs) raises several ethical concerns. In creating an open model, we have carefully considered the following:
+
+* **Bias and Fairness**
+ * VLMs trained on large-scale, real-world text and image data can reflect socio-cultural biases embedded in the training material. Gemma 4 models underwent careful scrutiny, input data pre-processing, and post-training evaluations as reported in this card to help mitigate the risk of these biases.
+* **Misinformation and Misuse**
+ * VLMs can be misused to generate text that is false, misleading, or harmful.
+ * Guidelines are provided for responsible use with the model, see the [Responsible Generative AI Toolkit](https://ai.google.dev/responsible).
+* **Transparency and Accountability**
+ * This model card summarizes details on the models' architecture, capabilities, limitations, and evaluation processes.
+ * A responsibly developed open model offers the opportunity to share innovation by making VLM technology accessible to developers and researchers across the AI ecosystem.
+
+**Risks identified and mitigations**:
-- [unsloth/gemma-4-12B-it-qat-GGUF](https://huggingface.co/unsloth/gemma-4-12B-it-qat-GGUF)
-- [Gemma 4 license](https://ai.google.dev/gemma/docs/gemma_4_license)
+* **Generation of harmful content**: Mechanisms and guidelines for content safety are essential. Developers are encouraged to exercise caution and implement appropriate content safety safeguards based on their specific product policies and application use cases.
+* **Misuse for malicious purposes**: Technical limitations and developer and end-user education can help mitigate against malicious applications of VLMs. Educational resources and reporting mechanisms for users to flag misuse are provided.
+* **Privacy violations**: Models were trained on data filtered for removal of certain personal information and other sensitive data. Developers are encouraged to adhere to privacy regulations with privacy-preserving techniques.
+* **Perpetuation of biases**: It's encouraged to perform continuous monitoring (using evaluation metrics, human review) and the exploration of de-biasing techniques during model training, fine-tuning, and other use cases.
-## Upload repo lên Hugging Face
+### **Benefits**
-Xem [UPLOAD.md](./UPLOAD.md) (trong project gốc: `huggingface/UPLOAD.md`).
+At the time of release, this family of models provides high-performance open vision-language model implementations designed from the ground up for responsible AI development compared to similarly sized models.
\ No newline at end of file
diff --git "a/Ti\341\273\203u \303\201i Test.mp4" "b/Ti\341\273\203u \303\201i Test.mp4"
new file mode 100644
index 0000000000000000000000000000000000000000..b13aa276784918afcdb3bc488cefba95b674fcc3
--- /dev/null
+++ "b/Ti\341\273\203u \303\201i Test.mp4"
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:af7d511664225d358167af3bb5163731de3838010f65e0814d10277ad5bc3efa
+size 659294678
diff --git "a/Ti\341\273\203u \303\201i Test.srt" "b/Ti\341\273\203u \303\201i Test.srt"
new file mode 100644
index 0000000000000000000000000000000000000000..9b9d1a17856ae7859789aebfc34f8b7c661cb222
--- /dev/null
+++ "b/Ti\341\273\203u \303\201i Test.srt"
@@ -0,0 +1,855 @@
+1
+00:00:10,261 --> 00:00:11,861
+不愧是小提琴女王
+
+2
+00:00:12,021 --> 00:00:13,301
+从相貌到家世
+
+3
+00:00:13,301 --> 00:00:14,581
+哪样都很出色
+
+4
+00:00:15,061 --> 00:00:16,821
+听说她是天赋型选手
+
+5
+00:00:17,141 --> 00:00:18,741
+天生对音乐很敏锐
+
+6
+00:00:18,901 --> 00:00:20,341
+学了很多乐器
+
+7
+00:00:20,341 --> 00:00:22,101
+却独独钟爱小提琴
+
+8
+00:00:48,501 --> 00:00:49,621
+02
+
+9
+00:00:56,821 --> 00:00:57,461
+柒柒
+
+10
+00:00:57,941 --> 00:00:58,741
+在忙吗
+
+11
+00:01:00,021 --> 00:01:00,661
+顾总
+
+12
+00:01:00,661 --> 00:01:03,061
+时老师在接受采访呢
+
+13
+00:01:03,061 --> 00:01:04,501
+需要我把手机给她吗
+
+14
+00:01:06,421 --> 00:01:07,221
+不用了
+
+15
+00:01:39,061 --> 00:01:39,701
+谢谢大家
+
+16
+00:01:39,861 --> 00:01:40,661
+那我们今天的采访
+
+17
+00:01:40,661 --> 00:01:41,461
+就到此结束了
+
+18
+00:01:41,461 --> 00:01:41,961
+谢谢
+
+19
+00:01:42,901 --> 00:01:44,181
+时小姐等一下
+
+20
+00:01:44,181 --> 00:01:45,781
+麻烦再回答一个问题好吗
+
+21
+00:01:45,781 --> 00:01:46,581
+不好意思
+
+22
+00:01:57,141 --> 00:01:58,741
+时老师
+
+23
+00:01:58,741 --> 00:01:59,381
+你真不愧是
+
+24
+00:01:59,381 --> 00:02:00,981
+名誉世界的小提琴女王
+
+25
+00:02:00,981 --> 00:02:02,581
+简直是太完美了
+
+26
+00:02:04,821 --> 00:02:05,461
+呆甜
+
+27
+00:02:06,101 --> 00:02:07,381
+擦擦你的口水
+
+28
+00:02:12,821 --> 00:02:13,941
+时老师
+
+29
+00:02:13,941 --> 00:02:15,381
+刚才顾总来电话了
+
+30
+00:02:17,781 --> 00:02:18,421
+知道了
+
+31
+00:02:18,901 --> 00:02:20,661
+顾家和时家门当户对
+
+32
+00:02:20,661 --> 00:02:21,461
+当年
+
+33
+00:02:21,461 --> 00:02:23,221
+时家融资出现问题了
+
+34
+00:02:23,221 --> 00:02:23,861
+是顾总出资
+
+35
+00:02:23,861 --> 00:02:25,461
+帮时家渡过难关的
+
+36
+00:02:25,621 --> 00:02:26,421
+婚后
+
+37
+00:02:26,421 --> 00:02:27,701
+两个人有孩子了
+
+38
+00:02:27,701 --> 00:02:28,981
+顾总对时老师
+
+39
+00:02:28,981 --> 00:02:29,941
+事无巨细
+
+40
+00:02:29,941 --> 00:02:30,901
+宠爱有加
+
+41
+00:02:35,701 --> 00:02:36,501
+时老师
+
+42
+00:02:36,821 --> 00:02:38,741
+那个小少爷生日快到了
+
+43
+00:02:38,741 --> 00:02:41,301
+需不需要准备点礼物寄回国
+
+44
+00:02:46,101 --> 00:02:46,741
+不需要
+
+45
+00:02:50,901 --> 00:02:51,701
+朝
+
+46
+00:02:55,145 --> 00:02:57,225
+就知道会是这样的答案
+
+47
+00:02:57,545 --> 00:02:58,665
+因为小少爷身主
+
+48
+00:02:58,665 --> 00:02:59,945
+流着顾总的血
+
+49
+00:03:00,425 --> 00:03:01,545
+所以时老师
+
+50
+00:03:01,545 --> 00:03:02,985
+连带小少爷
+
+51
+00:03:03,465 --> 00:03:04,585
+也很冷淡
+
+52
+00:03:06,185 --> 00:03:06,985
+等下个月回国了
+
+53
+00:03:06,985 --> 00:03:08,105
+咱们好好庆祝一下
+
+54
+00:03:12,265 --> 00:03:12,905
+哎
+
+55
+00:03:13,065 --> 00:03:13,865
+怎么了
+
+56
+00:03:15,145 --> 00:03:16,425
+脸色这么难看
+
+57
+00:03:18,185 --> 00:03:18,985
+这么烫
+
+58
+00:03:19,305 --> 00:03:20,425
+你先带柒柒去医院
+
+59
+00:03:20,425 --> 00:03:21,225
+这边我来收尾
+
+60
+00:03:21,225 --> 00:03:21,725
+好
+
+61
+00:03:24,425 --> 00:03:25,865
+时老师你没事吧
+
+62
+00:03:28,745 --> 00:03:29,705
+这位患者
+
+63
+00:03:29,705 --> 00:03:31,305
+是急性肠胃炎导致的高烧
+
+64
+00:03:32,265 --> 00:03:32,765
+注意观察
+
+65
+00:03:36,425 --> 00:03:37,225
+小宝乖
+
+66
+00:03:38,665 --> 00:03:39,625
+明天做完手术
+
+67
+00:03:39,625 --> 00:03:40,905
+我们就可以回家了
+
+68
+00:03:41,385 --> 00:03:43,305
+到时候妈妈再带你去游乐场
+
+69
+00:03:43,305 --> 00:03:44,585
+补过生日好不好
+
+70
+00:03:44,905 --> 00:03:47,305
+妈妈那我们说好了
+
+71
+00:03:47,465 --> 00:03:49,225
+明天做完手术后
+
+72
+00:03:49,225 --> 00:03:51,305
+就带我去过生日
+
+73
+00:03:52,745 --> 00:03:53,225
+好
+
+74
+00:03:53,225 --> 00:03:54,185
+妈妈答应你
+
+75
+00:03:59,785 --> 00:04:00,265
+哎
+
+76
+00:04:00,265 --> 00:04:01,065
+真可怜啊
+
+77
+00:04:01,545 --> 00:04:02,985
+手术费还差十万元呢
+
+78
+00:04:03,945 --> 00:04:05,545
+明天就算筹齐善款
+
+79
+00:04:05,545 --> 00:04:06,345
+可以做手术
+
+80
+00:04:06,665 --> 00:04:08,425
+手术成功的可能也挺小
+
+81
+00:04:09,705 --> 00:04:11,785
+孩子才四岁就患白血病了
+
+82
+00:04:11,785 --> 00:04:12,745
+真是可怜
+
+83
+00:04:13,225 --> 00:04:14,665
+但他妈妈真的爱他儿子
+
+84
+00:04:14,665 --> 00:04:16,745
+每天给他陪床做饭
+
+85
+00:04:17,225 --> 00:04:17,705
+妈妈
+
+86
+00:04:17,705 --> 00:04:18,205
+别哭了
+
+87
+00:04:24,905 --> 00:04:25,545
+小姐
+
+88
+00:04:25,705 --> 00:04:27,305
+你真的愿意帮那对母子
+
+89
+00:04:27,305 --> 00:04:28,105
+补交费用
+
+90
+00:04:30,025 --> 00:04:30,665
+不过
+
+91
+00:04:30,985 --> 00:04:33,065
+你不用告诉他们是谁给的
+
+92
+00:04:34,185 --> 00:04:34,985
+就说
+
+93
+00:04:35,465 --> 00:04:36,905
+是小宝的生日礼物
+
+94
+00:04:39,465 --> 00:04:41,545
+真是有钱有颜还善良
+
+95
+00:04:45,225 --> 00:04:46,505
+订最近的一个航班
+
+96
+00:04:47,145 --> 00:04:47,645
+回国
+
+97
+00:04:47,785 --> 00:04:48,285
+哎
+
+98
+00:04:48,585 --> 00:04:49,545
+时老师
+
+99
+00:04:49,705 --> 00:04:51,145
+你这身体还吃得消吗
+
+100
+00:04:51,625 --> 00:04:52,425
+要不然
+
+101
+00:04:52,425 --> 00:04:54,345
+在医院住两天再走吧
+
+102
+00:04:55,465 --> 00:04:56,105
+妈妈
+
+103
+00:04:56,105 --> 00:04:57,545
+别哭了
+
+104
+00:04:59,145 --> 00:04:59,645
+不用
+
+105
+00:05:00,425 --> 00:05:01,225
+订机票吧
+
+106
+00:05:07,465 --> 00:05:08,105
+师傅
+
+107
+00:05:08,585 --> 00:05:09,705
+去星河别墅
+
+108
+00:05:10,185 --> 00:05:10,685
+好的
+
+109
+00:05:13,225 --> 00:05:14,505
+这好像
+
+110
+00:05:14,505 --> 00:05:16,905
+是我跟顾寒深结婚五年来
+
+111
+00:05:17,385 --> 00:05:18,505
+第一次
+
+112
+00:05:18,505 --> 00:05:20,425
+演出一结束就回家
+
+113
+00:05:21,225 --> 00:05:22,665
+你不爱我没关系
+
+114
+00:05:23,145 --> 00:05:24,905
+哪怕恨我都没关系
+
+115
+00:05:27,625 --> 00:05:28,585
+五年前
+
+116
+00:05:28,905 --> 00:05:30,985
+时家资金链断裂
+
+117
+00:05:30,985 --> 00:05:32,105
+面临破产
+
+118
+00:05:33,225 --> 00:05:35,145
+顾寒深伸出援手
+
+119
+00:05:36,265 --> 00:05:37,545
+条件却是
+
+120
+00:05:38,985 --> 00:05:40,745
+只要肯做我太太就好
+
+121
+00:05:43,145 --> 00:05:44,585
+要我跟他结婚
+
+122
+00:05:51,354 --> 00:05:52,794
+你不爱我没关系
+
+123
+00:05:54,234 --> 00:05:55,834
+哪怕恨我都没关系
+
+124
+00:05:57,274 --> 00:05:59,034
+只要肯做我太太就好
+
+125
+00:06:07,514 --> 00:06:08,634
+只要你点头
+
+126
+00:06:09,594 --> 00:06:11,194
+我立刻让人把五个亿
+
+127
+00:06:11,994 --> 00:06:13,434
+打入你们时氏的账户
+
+128
+00:06:16,474 --> 00:06:17,754
+我憎恨他做局
+
+129
+00:06:18,394 --> 00:06:19,834
+害我当时的男朋友
+
+130
+00:06:19,834 --> 00:06:21,114
+沈景泽入狱
+
+131
+00:06:24,154 --> 00:06:25,594
+以至于结婚后
+
+132
+00:06:31,194 --> 00:06:32,634
+对他恨之入骨
+
+133
+00:06:39,674 --> 00:06:40,314
+柒柒
+
+134
+00:06:40,794 --> 00:06:41,594
+我求你
+
+135
+00:06:41,594 --> 00:06:42,394
+你别走
+
+136
+00:06:43,354 --> 00:06:44,154
+别碰我
+
+137
+00:06:44,634 --> 00:06:45,434
+滚
+
+138
+00:07:20,474 --> 00:07:22,394
+只要你肯生下这个孩子
+
+139
+00:07:29,594 --> 00:07:31,034
+我就放你自由
+
+140
+00:07:35,994 --> 00:07:37,594
+直到上个月我才知道
+
+141
+00:07:38,234 --> 00:07:39,674
+当初转移股份
+
+142
+00:07:40,474 --> 00:07:41,754
+截断资金链
+
+143
+00:07:41,754 --> 00:07:43,514
+让公司陷入危机的人
+
+144
+00:07:47,034 --> 00:07:48,634
+是我引狼入室
+
+145
+00:07:48,954 --> 00:07:50,234
+看错了人
+
+146
+00:07:50,874 --> 00:07:51,674
+柒柒
+
+147
+00:07:51,994 --> 00:07:54,234
+你让我调查的事情有结果了
+
+148
+00:07:55,194 --> 00:07:55,834
+当年
+
+149
+00:07:55,834 --> 00:07:57,274
+不是顾寒深设的局
+
+150
+00:07:58,074 --> 00:07:59,034
+你误会他了
+
+151
+00:07:59,514 --> 00:08:01,274
+当年如果没有他出资相助
+
+152
+00:08:01,594 --> 00:08:02,714
+时家现在
+
+153
+00:08:02,714 --> 00:08:03,834
+恐怕早就没了
+
+154
+00:08:06,554 --> 00:08:07,674
+而我这么多年
+
+155
+00:08:08,314 --> 00:08:10,074
+一直都在错怪顾寒深
+
+156
+00:08:10,874 --> 00:08:11,674
+可他
+
+157
+00:08:12,314 --> 00:08:13,594
+却无一怨言
+
+158
+00:08:24,314 --> 00:08:26,074
+小姐星河别墅到了
+
+159
+00:08:30,874 --> 00:08:31,674
+谢谢
+
+160
+00:08:36,954 --> 00:08:38,714
+因为我喜欢玫瑰
+
+161
+00:08:39,354 --> 00:08:40,954
+所以顾寒深
+
+162
+00:08:41,274 --> 00:08:42,554
+持巨资
+
+163
+00:08:43,194 --> 00:08:45,114
+为我打造了这片玫瑰花海
+
+164
+00:08:51,097 --> 00:08:53,177
+是棠棠告诉我真相的原因吗
+
+165
+00:08:54,457 --> 00:08:55,577
+这次回来
+
+166
+00:08:56,057 --> 00:08:57,017
+竟然对这个
+
+167
+00:08:57,017 --> 00:08:58,777
+曾经厌恶至极的家
+
+168
+00:08:59,257 --> 00:09:01,337
+产生了一丝丝的思念
+
+169
+00:09:03,417 --> 00:09:04,217
+太太
+
+170
+00:09:05,337 --> 00:09:06,137
+太太
+
+171
+00:09:06,297 --> 00:09:06,937
+您回来了
+
+172
+00:09:06,937 --> 00:09:08,697
+怎么没有提前跟我说一声
+
+173
+00:09:08,697 --> 00:09:10,457
+我好派车去接你啊
+
+174
+00:09:11,257 --> 00:09:12,857
+先生他知道了吗
+
+175
+00:09:15,897 --> 00:09:16,697
+他还不知道
+
+176
+00:09:17,657 --> 00:09:18,617
+那快进屋吧
+
+177
+00:09:25,497 --> 00:09:26,297
+暮暮
+
+178
+00:09:26,937 --> 00:09:28,057
+已经十点了
+
+179
+00:09:28,377 --> 00:09:29,657
+明天还要上学
+
+180
+00:09:29,817 --> 00:09:30,777
+快去睡觉
+
+181
+00:09:31,257 --> 00:09:32,217
+爸爸
+
+182
+00:09:32,537 --> 00:09:35,097
+妈妈怎么还没有回来
+
+183
+00:09:35,577 --> 00:09:39,417
+她走之前答应给我过生日的
+
+184
+00:09:40,217 --> 00:09:43,097
+你不是也说她会回来的吗
+
+185
+00:09:43,097 --> 00:09:44,377
+妈妈是
+
+186
+00:09:45,017 --> 00:09:47,897
+是不是不喜欢我
+
+187
+00:09:53,817 --> 00:09:54,937
+怎么会呢
+
+188
+00:09:55,577 --> 00:09:57,177
+妈妈很喜欢暮暮
+
+189
+00:09:58,617 --> 00:10:00,217
+妈妈工作很忙
+
+190
+00:10:00,857 --> 00:10:02,617
+没有办法能回来陪你
+
+191
+00:10:03,097 --> 00:10:05,177
+明天爸爸送暮暮去上幼儿园
+
+192
+00:10:05,177 --> 00:10:05,977
+好不好
+
+193
+00:10:06,617 --> 00:10:08,057
+你骗我
+
+194
+00:10:08,217 --> 00:10:10,617
+妈妈就是不喜欢我
+
+195
+00:10:10,777 --> 00:10:12,857
+别人家的小朋友
+
+196
+00:10:12,857 --> 00:10:15,897
+妈妈都会给他们过生日
+
+197
+00:10:16,057 --> 00:10:17,817
+送他们上幼儿园
+
+198
+00:10:17,977 --> 00:10:19,897
+给他们讲故事
+
+199
+00:10:20,057 --> 00:10:24,537
+我妈妈从来不理我
+
+200
+00:10:33,177 --> 00:10:34,137
+是啊
+
+201
+00:10:34,777 --> 00:10:36,697
+暮暮都看得清的事实
+
+202
+00:10:37,497 --> 00:10:39,577
+我却一直在自欺欺人
+
+203
+00:10:41,017 --> 00:10:42,137
+她不喜欢我
+
+204
+00:10:42,617 --> 00:10:44,377
+也不喜欢我们的孩子
+
+205
+00:10:47,417 --> 00:10:49,817
+太太本就不喜欢小少爷
+
+206
+00:10:49,817 --> 00:10:51,577
+小少爷这么哭下去
+
+207
+00:10:51,577 --> 00:10:53,657
+会不会哭恼了太太
+
+208
+00:10:55,097 --> 00:10:55,897
+太太
+
+209
+00:10:55,897 --> 00:10:57,497
+您怎么不进去啊
+
+210
+00:11:12,697 --> 00:11:13,977
+柒柒
+
+211
+00:11:44,057 --> 00:11:45,177
+你抱够了没有
+
+212
+00:11:46,937 --> 00:11:49,177
+朝思暮时
+
+213
+00:11:58,674 --> 00:11:59,634
+这些年
+
+214
+00:11:59,634 --> 00:12:00,754
+我每次回家
\ No newline at end of file
diff --git a/__pycache__/diarize_audio.cpython-310.pyc b/__pycache__/diarize_audio.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..d0b80021d8a58e69af887da2830424c663dccc09
Binary files /dev/null and b/__pycache__/diarize_audio.cpython-310.pyc differ
diff --git a/__pycache__/translate_srt.cpython-310.pyc b/__pycache__/translate_srt.cpython-310.pyc
new file mode 100644
index 0000000000000000000000000000000000000000..83613631b08fc9f754a3772f3b55ed1254703222
Binary files /dev/null and b/__pycache__/translate_srt.cpython-310.pyc differ
diff --git a/hf-upload/README.md b/hf-upload/README.md
new file mode 100644
index 0000000000000000000000000000000000000000..bcdab3d85d5123ce8157b684fae3c63d1ebb60b9
--- /dev/null
+++ b/hf-upload/README.md
@@ -0,0 +1,139 @@
+---
+license: apache-2.0
+language:
+ - vi
+ - en
+tags:
+ - subtitles
+ - srt
+ - translation
+ - gemma4
+ - vision
+ - llama.cpp
+pipeline_tag: translation
+library_name: gemma-srt-translate
+---
+
+# Gemma SRT Translate
+
+Dịch và sửa phụ đề SRT bằng **Gemma 4 12B vision** (llama-server + MTP).
+Repo này chứa **script + notebook Colab**; model GGUF tải từ [`unsloth/gemma-4-12B-it-qat-GGUF`](https://huggingface.co/unsloth/gemma-4-12B-it-qat-GGUF).
+
+## Cấu trúc repo
+
+```
+YOUR_USERNAME/gemma-srt-translate/
+├── README.md ← file này
+├── config.yaml ← cấu hình mặc định
+├── translate_srt.py ← pipeline chính
+├── diarize_audio.py ← Pass 0: phân tích giọng nói (tùy chọn)
+├── requirements-diarize.txt ← dependency cho Pass 0
+├── colab/
+│ └── GemmaSRT_Colab.ipynb ← chạy trên Google Colab (L4 24GB)
+└── scripts/
+ ├── download_models.py ← tải GGUF từ Unsloth
+ └── build_llama_server.sh ← build llama-server Linux (Colab)
+```
+
+## Chạy trên Google Colab (khuyến nghị L4 24GB)
+
+Repo **Private** → **không** mở được bằng nút Open in Colab trên HF (lỗi 401).
+
+**Cách mở:** tải [colab/GemmaSRT_Colab.ipynb](https://huggingface.co/STBack23/gemma-srt-translate/blob/main/colab/GemmaSRT_Colab.ipynb) → Colab **File → Upload notebook** → dán **`HF_TOKEN`** vào cell cấu hình → lưu copy lên Drive `Gemma/` để dùng lại.
+
+Repo này nằm trên **Hugging Face**, không phải GitHub — link `colab.research.google.com/github/...` sẽ lỗi 404.
+
+**Cách mở (chọn một):**
+
+1. **Nút Open in Colab trên HF** (khuyến nghị):
+ [colab/GemmaSRT_Colab.ipynb](https://huggingface.co/STBack23/gemma-srt-translate/blob/main/colab/GemmaSRT_Colab.ipynb) → bấm **Open in Colab**
+
+2. **Shortcut HF `/colab`**:
+ [https://huggingface.co/STBack23/gemma-srt-translate/colab](https://huggingface.co/STBack23/gemma-srt-translate/colab)
+
+3. **Colab → File → Upload notebook** → tải file `.ipynb` từ HF về rồi upload
+
+Sau khi mở notebook:
+
+1. Runtime → **Change runtime type** → GPU **L4** (hoặc T4/A100)
+2. Chạy tuần tự các cell (lần đầu ~15–20 phút: build llama-server + tải model)
+3. Upload video + SRT hoặc trỏ vào Google Drive
+4. Tải file `*.vi.srt` về máy
+
+[](https://huggingface.co/STBack23/gemma-srt-translate/colab)
+
+## Chạy trên máy local (Windows)
+
+Repo HF **không** chứa `llama-server.exe` — dùng bản Windows trong project gốc:
+
+```powershell
+.\run-translate.ps1 -Video "phim.mp4" -InputSrt "phim.srt" -OutputSrt "phim.vi.srt"
+```
+
+Hoặc mở GUI: `GemmaSRT.bat`
+
+## Tải model thủ công
+
+```bash
+pip install huggingface_hub
+python scripts/download_models.py --dest ./models
+```
+
+File tải về (~12–15 GB):
+
+| File | Mục đích |
+|------|----------|
+| `gemma-4-12B-it-qat-UD-Q4_K_XL.gguf` | Model chính |
+| `mmproj-F16.gguf` | Vision projector |
+| `mtp-gemma-4-12B-it.gguf` | MTP draft (tăng tốc) |
+
+## Pipeline
+
+1. Parse SRT → gom scene
+2. **Pass 0** *(tùy chọn)*: phân tích giọng nói (pyannote) → "ai nói câu nào" + giới tính/độ tuổi → gán vào từng cue để xưng hô nhất quán
+3. **Pass 1**: cắt frame (ffmpeg) → OCR/sửa SRT gốc (vision)
+4. **Pass 2**: dịch theo ngữ cảnh + xưng hô
+5. Ghi SRT dịch, SRT đã sửa, báo cáo JSON
+
+### Pass 0 — phân tích giọng nói (diarization)
+
+Bật bằng `--diarize` (CLI) hoặc `ENABLE_DIARIZE = True` (Colab). Cần thêm:
+
+```bash
+pip install -r requirements-diarize.txt
+```
+
+Token HF phải bấm **Agree** điều kiện các model gated:
+[community-1](https://hf.co/pyannote/speaker-diarization-community-1) ·
+[3.1](https://hf.co/pyannote/speaker-diarization-3.1) ·
+[segmentation-3.0](https://hf.co/pyannote/segmentation-3.0).
+`--gender-method model` dùng [audeering wav2vec2 age/gender](https://hf.co/audeering/wav2vec2-large-robust-24-ft-age-gender) (~1GB).
+
+## VRAM
+
+| GPU | 1 phim | 2 phim song song |
+|-----|--------|------------------|
+| RTX 4060 Ti 16GB | ✅ | ❌ |
+| Colab L4 24GB | ✅ | ❌ (dùng local + Colab = 2 phim độc lập) |
+
+## Cấu hình
+
+Chỉnh `config.yaml` hoặc tham số CLI:
+
+```bash
+python translate_srt.py \
+ --video phim.mp4 \
+ --input-srt phim.srt \
+ --output-srt phim.vi.srt \
+ --target-lang Vietnamese \
+ --skip-correction # bỏ pass OCR (nhanh hơn)
+```
+
+## Model gốc
+
+- [unsloth/gemma-4-12B-it-qat-GGUF](https://huggingface.co/unsloth/gemma-4-12B-it-qat-GGUF)
+- [Gemma 4 license](https://ai.google.dev/gemma/docs/gemma_4_license)
+
+## Upload repo lên Hugging Face
+
+Xem [UPLOAD.md](./UPLOAD.md) (trong project gốc: `huggingface/UPLOAD.md`).
diff --git a/hf-upload/colab/GemmaSRT_Colab.ipynb b/hf-upload/colab/GemmaSRT_Colab.ipynb
new file mode 100644
index 0000000000000000000000000000000000000000..5998bb6e1a7c2b9418f5911c8ef4da74d9838ddb
--- /dev/null
+++ b/hf-upload/colab/GemmaSRT_Colab.ipynb
@@ -0,0 +1,633 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {},
+ "source": [
+ "# Gemma SRT Translate — Google Colab\n",
+ "\n",
+ "Dịch phụ đề `.srt` sang tiếng Việt bằng **Gemma 4 12B vision**.\n",
+ "\n",
+ "**GPU:** T4/L4 — Runtime → Change runtime type → GPU. \n",
+ "\n",
+ "## Chạy\n",
+ "\n",
+ "Lần đầu ~20–50 phút (tải model ~12GB + build llama-server; cache lưu `Gemma/Cache`).\n",
+ "\n",
+ "## Tuỳ chọn (cell CẤU HÌNH)\n",
+ "\n",
+ "- `ENABLE_DIARIZE = True` — phân tích giọng nói để xưng hô nhất quán (cần video)\n",
+ "- `SKIP_CORRECTION = True` — bỏ OCR sửa SRT. Đặt `False` nếu sub gốc hay sai chữ\n",
+ "- `FORCE_LLAMA_REBUILD = True` — build lại llama-server b9553 (MTP); đặt `False` sau lần build OK\n",
+ "- `CTX_SIZE = 4096` — context nhỏ hơn giúp MTP chạy trên L4\n",
+ "- `NO_SCENE_VISION` / `SCENE_FRAMES` — tắt/giảm vision để dịch nhanh hơn (phim dài)\n",
+ "- `DOWNLOAD_RESULTS` — tải `.vi.srt` về trình duyệt sau khi xong"
+ ]
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# ═══ CẤU HÌNH — sửa trước khi chạy ═══\n",
+ "from pathlib import Path\n",
+ "import os\n",
+ "\n",
+ "from google.colab import userdata\n",
+ "\n",
+ "HF_TOKEN_SECRET = \"HF_TOKEN\" # tên khóa trong Colab Secrets (🔑)\n",
+ "\n",
+ "def _load_hf_token() -> str:\n",
+ " try:\n",
+ " return userdata.get(HF_TOKEN_SECRET).strip()\n",
+ " except userdata.SecretNotFoundError:\n",
+ " return \"\"\n",
+ "\n",
+ "HF_TOKEN = _load_hf_token()\n",
+ "if HF_TOKEN:\n",
+ " os.environ[\"HF_TOKEN\"] = HF_TOKEN\n",
+ " os.environ[\"HUGGING_FACE_HUB_TOKEN\"] = HF_TOKEN\n",
+ " print(f\"HF token OK (Secret: {HF_TOKEN_SECRET})\")\n",
+ "else:\n",
+ " print(f\"Chưa có Secret '{HF_TOKEN_SECRET}' — thêm tại 🔑 Secrets (Read token HF)\")\n",
+ "\n",
+ "HF_REPO = \"STBack23/gemma-srt-translate\" # repo chứa mã nguồn pipeline\n",
+ "\n",
+ "MODELS_REPO = \"unsloth/gemma-4-12B-it-qat-GGUF\" # ~12GB — T4/L4\n",
+ "MODEL_FILES = [\n",
+ " \"gemma-4-12B-it-qat-UD-Q4_K_XL.gguf\",\n",
+ " \"mmproj-F16.gguf\",\n",
+ " \"mtp-gemma-4-12B-it.gguf\",\n",
+ "]\n",
+ "\n",
+ "# ── Google Drive ─────────────────────────────────────────────\n",
+ "USE_DRIVE_CACHE = True # True = cache model + llama-server trên Drive (khuyến nghị)\n",
+ "\n",
+ "# Thư mục trên Drive: Gemma/Cache (model) + Gemma/Phim (video + SRT vào/ra)\n",
+ "DRIVE_ROOT = \"/content/drive/MyDrive/Gemma\"\n",
+ "DRIVE_CACHE = f\"{DRIVE_ROOT}/Cache\"\n",
+ "DRIVE_PHIM_DIR = f\"{DRIVE_ROOT}/Phim\"\n",
+ "\n",
+ "# Mã nguồn pipeline luôn tải từ Hugging Face (HF_REPO).\n",
+ "VIDEO_EXTS = (\".mp4\", \".mkv\", \".avi\", \".mov\")\n",
+ "\n",
+ "def job(stem, ext_video=\".mp4\"):\n",
+ " \"\"\"stem = tên file trên Drive (KHÔNG đuôi), khớp y hệt — kể cả dấu, khoảng trắng.\"\"\"\n",
+ " base = f\"{DRIVE_PHIM_DIR}/{stem}\"\n",
+ " return {\"stem\": stem, \"video\": f\"{base}{ext_video}\", \"srt\": f\"{base}.srt\"}\n",
+ "\n",
+ "# ── Danh sách phim cần dịch ──────────────────────────────────\n",
+ "PHIM_STEMS = [\n",
+ " \"Thư Gửi Mùa Hè\", # → Gemma/Phim/Thư Gửi Mùa Hè.mp4 + Thư Gửi Mùa Hè.srt\n",
+ " # \"ten-phim-khac\",\n",
+ "]\n",
+ "JOBS = [job(stem) for stem in PHIM_STEMS]\n",
+ "# Cách khác — đường dẫn đầy đủ:\n",
+ "# JOBS = [{\"video\": f\"{DRIVE_PHIM_DIR}/a.mp4\", \"srt\": f\"{DRIVE_PHIM_DIR}/a.srt\"}]\n",
+ "# Upload 1 phim qua nút Colab: PHIM_STEMS = [], JOBS = [], UPLOAD_WIDGET = True\n",
+ "UPLOAD_WIDGET = False\n",
+ "\n",
+ "# ── Tùy chọn dịch ───────────────────────────────────────────\n",
+ "SOURCE_LANG = \"auto\" # ngôn ngữ SRT gốc (\"auto\" = model tự nhận)\n",
+ "TARGET_LANG = \"Vietnamese\" # ngôn ngữ đích\n",
+ "\n",
+ "SKIP_CORRECTION = True # False = OCR sửa SRT trước khi dịch\n",
+ "\n",
+ "LIMIT_CUES = 0 # 0 = dịch hết file; đặt 20 để thử nhanh vài cue đầu\n",
+ "\n",
+ "ENABLE_DIARIZE = True # False = bỏ phân tích giọng nói\n",
+ "\n",
+ "GENDER_METHOD = \"model\" # \"model\" | \"pitch\" | \"auto\"\n",
+ "\n",
+ "NUM_SPEAKERS = 0 # 0 = tự dò; đặt số nếu biết trước số nhân vật chính\n",
+ "DETECT_GENDER = True # False = chỉ gán nhãn người nói, bỏ giới tính/tuổi\n",
+ "DIARIZE_DEVICE = \"auto\" # \"auto\" | \"cuda\" | \"cpu\"\n",
+ "\n",
+ "# ── llama-server + MTP (Gemma 4 draft-mtp) ───────────────────\n",
+ "# Pin llama.cpp tag b9553 — verified Gemma 4 draft-mtp on L4.\n",
+ "FORCE_LLAMA_REBUILD = True # True = xóa cache Drive + build lại (~20–50 phút)\n",
+ " # Sau khi log có \"(with MTP)\", đặt False cho lần sau\n",
+ "NO_MTP = False # True = tắt MTP (chậm hơn ~20–40%)\n",
+ "CTX_SIZE = 4096 # 4096 ổn L4+MTP; 8192 nếu đủ VRAM / không MTP\n",
+ "\n",
+ "# ── Tốc độ dịch (phim dài) ───────────────────────────────────\n",
+ "NO_SCENE_VISION = False # True = bỏ bước mô tả cảnh (~40% nhanh hơn)\n",
+ "NO_SCENE_IMAGE = False # True = dịch không ảnh (nhanh nhất, kém context)\n",
+ "SCENE_FRAMES = 3 # 1 = ít ảnh/scene, nhanh hơn\n",
+ "SCENE_MAX_CUES = 4 # 5–6 = ít lần gọi model hơn (phim dài)\n",
+ "SCENE_MAX_GAP = 1.5 # giây — tách scene nếu khoảng cách lớn hơn\n",
+ "SCENE_MAX_DUR = 20.0 # giây — tách scene nếu đoạn quá dài\n",
+ "\n",
+ "# ── Sau khi dịch xong ─────────────────────────────────────────\n",
+ "DOWNLOAD_RESULTS = True # True = tải *.vi.srt về trình duyệt\n",
+ "AUTO_UNMOUNT_DRIVE = True # True = gỡ mount Google Drive\n",
+ "AUTO_DISCONNECT_RUNTIME = True # True = ngắt runtime Colab (tiết kiệm GPU quota)\n",
+ "DISCONNECT_DELAY_SEC = 15 # giây chờ trước khi ngắt (để download kịp)\n",
+ "\n",
+ "# ── Đường dẫn nội bộ Colab (thường không cần sửa) ────────────\n",
+ "ROOT = \"/content/gemma-srt\"\n",
+ "MODELS_DIR = f\"{ROOT}/models\"\n",
+ "LLAMA_DIR = \"/content/llama.cpp\"\n",
+ "LLAMA_SERVER = f\"{LLAMA_DIR}/build/bin/llama-server\""
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# Kiểm tra GPU\n",
+ "!nvidia-smi --query-gpu=name,memory.total --format=csv,noheader\n",
+ "\n",
+ "import torch\n",
+ "if not torch.cuda.is_available():\n",
+ " raise RuntimeError(\"Chưa có GPU. Runtime → Change runtime type → GPU (L4).\")"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# Mount Google Drive — cache model + thư mục phim\n",
+ "if USE_DRIVE_CACHE:\n",
+ " from google.colab import drive\n",
+ " from pathlib import Path\n",
+ " import os\n",
+ " drive.mount(\"/content/drive\")\n",
+ " os.makedirs(DRIVE_CACHE, exist_ok=True)\n",
+ " os.makedirs(DRIVE_PHIM_DIR, exist_ok=True)\n",
+ " print(f\"Cache model: {DRIVE_CACHE}\")\n",
+ " print(f\"Thư mục phim: {DRIVE_PHIM_DIR}\")\n",
+ "\n",
+ " cache_tgz = Path(f\"{DRIVE_CACHE}/llama-server-bin.tgz\")\n",
+ " if cache_tgz.is_file():\n",
+ " mb = cache_tgz.stat().st_size / 1_048_576\n",
+ " print(f\" llama-server cache: {mb:.1f} MB\" + (\" OK\" if mb > 5 else \" HỎNG\"))\n",
+ " else:\n",
+ " print(\" llama-server cache: chưa có (lần đầu build ~20–50 phút, pin b9553 cho MTP)\")\n",
+ "\n",
+ " files = sorted(p.name for p in Path(DRIVE_PHIM_DIR).iterdir() if p.is_file())\n",
+ " if files:\n",
+ " print(f\" ({len(files)} file trong Phim)\")\n",
+ " for name in files[:12]:\n",
+ " print(f\" - {name}\")\n",
+ " if len(files) > 12:\n",
+ " print(f\" ... và {len(files) - 12} file khác\")\n",
+ " else:\n",
+ " print(\" (chưa có file — upload video + .srt vào Gemma/Phim)\")"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# Chuẩn bị code pipeline — tải từ Hugging Face\n",
+ "!pip install -q huggingface_hub pyyaml\n",
+ "\n",
+ "from pathlib import Path\n",
+ "import os\n",
+ "from huggingface_hub import snapshot_download\n",
+ "\n",
+ "Path(ROOT).mkdir(parents=True, exist_ok=True)\n",
+ "\n",
+ "_hf_tok = HF_TOKEN.strip()\n",
+ "if not _hf_tok or HF_REPO.startswith(\"YOUR_\"):\n",
+ " raise ValueError(\n",
+ " f\"Tải code từ HF cần HF_REPO hợp lệ + Secret '{HF_TOKEN_SECRET}' (Read token) — thêm tại 🔑 Secrets.\"\n",
+ " )\n",
+ "os.environ[\"HF_TOKEN\"] = _hf_tok\n",
+ "os.environ[\"HUGGING_FACE_HUB_TOKEN\"] = _hf_tok\n",
+ "\n",
+ "print(f\"[code] Downloading {HF_REPO} ...\")\n",
+ "snapshot_download(\n",
+ " repo_id=HF_REPO,\n",
+ " repo_type=\"model\",\n",
+ " local_dir=ROOT,\n",
+ " token=_hf_tok,\n",
+ " local_dir_use_symlinks=False,\n",
+ ")\n",
+ "print(f\"[code] OK (HF): {ROOT}\")\n",
+ "\n",
+ "assert Path(f\"{ROOT}/translate_srt.py\").is_file(), \"translate_srt.py not found\"\n",
+ "assert Path(f\"{ROOT}/scripts/ensure_llama_colab.py\").is_file(), \"scripts/ensure_llama_colab.py not found\""
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# ═══ Pass 0 — cài dependency phân tích giọng nói (chỉ khi ENABLE_DIARIZE) ═══\n",
+ "import importlib\n",
+ "import os\n",
+ "\n",
+ "if ENABLE_DIARIZE:\n",
+ " pip = get_ipython().run_line_magic\n",
+ " print(\"[diarize] Cài pyannote (Colab: giữ numpy<2.1 cho numba/scipy)...\")\n",
+ " # pyannote hay nâng numpy 2.4 → lỗi: cannot import name '_center' from numpy._core.umath\n",
+ " pip(\"pip\", \"install -q 'numpy>=1.26,<2.1'\")\n",
+ " pip(\"pip\", \"install -q 'lightning-utilities>=0.10' 'lightning>=2.2,<2.5'\")\n",
+ " pip(\"pip\", \"install -q pyannote.audio>=4.0 soundfile>=0.11 --upgrade-strategy only-if-needed\")\n",
+ " pip(\"pip\", \"install -q 'numpy>=1.26,<2.1' --force-reinstall --no-deps\")\n",
+ " try:\n",
+ " import transformers # noqa: F401\n",
+ " except Exception:\n",
+ " pip(\"pip\", \"install -q transformers>=4.40\")\n",
+ " importlib.invalidate_caches()\n",
+ " if HF_TOKEN.strip():\n",
+ " os.environ[\"HF_TOKEN\"] = HF_TOKEN.strip()\n",
+ " assert os.path.isfile(f\"{ROOT}/diarize_audio.py\"), \\\n",
+ " \"Thiếu diarize_audio.py — chạy lại cell tải code (Drive/HF) phía trên.\"\n",
+ " try:\n",
+ " import numpy as np\n",
+ " from pyannote.audio import Pipeline # noqa: F401\n",
+ " print(f\"[diarize] OK — numpy {np.__version__}, pyannote import OK.\")\n",
+ " print(\"[diarize] Trước khi chạy: tài khoản HF (Secret HF_TOKEN) phải bấm Agree tại:\")\n",
+ " for _u in (\n",
+ " \"https://hf.co/pyannote/speaker-diarization-community-1\",\n",
+ " \"https://hf.co/pyannote/speaker-diarization-3.1\",\n",
+ " \"https://hf.co/pyannote/segmentation-3.0\",\n",
+ " ):\n",
+ " print(f\" {_u}\")\n",
+ " except ImportError as e:\n",
+ " raise RuntimeError(\n",
+ " f\"pyannote import lỗi: {e}\\n\"\n",
+ " \"→ Runtime → Restart session → Run all lại (cell này chạy lại sau restart).\"\n",
+ " ) from e\n",
+ " if GENDER_METHOD == \"model\":\n",
+ " print(\"[diarize] Preflight model age/gender (xác nhận load + predict trước khi chạy)...\")\n",
+ " try:\n",
+ " import sys as _sys\n",
+ " if ROOT not in _sys.path:\n",
+ " _sys.path.insert(0, ROOT)\n",
+ " import numpy as _np\n",
+ " import torch as _torch\n",
+ " from diarize_audio import _load_age_gender_model, DEFAULT_AGE_GENDER_MODEL\n",
+ " _dev = \"cuda\" if _torch.cuda.is_available() else \"cpu\"\n",
+ " _m, _proc, _ = _load_age_gender_model(DEFAULT_AGE_GENDER_MODEL, _dev, print)\n",
+ " _sig = _np.zeros(16000, dtype=\"float32\")\n",
+ " _in = _proc(_sig, sampling_rate=16000)\n",
+ " _vals = _torch.from_numpy(_in[\"input_values\"][0].reshape(1, -1)).to(_torch.device(_dev))\n",
+ " with _torch.no_grad():\n",
+ " _h, _age, _gen = _m(_vals)\n",
+ " print(f\"[diarize] Preflight OK — model age/gender chạy được trên {_dev} (age head demo={float(_age[0].item())*100:.0f}y).\")\n",
+ " except Exception as _e: # noqa: BLE001\n",
+ " print(f\"[diarize] CANH BAO - Preflight age/gender LOI: {_e}\")\n",
+ " print(\"[diarize] -> Pass 0 van chay nhung FALLBACK pitch (chi gioi tinh, KHONG co tuoi).\")\n",
+ " print(\"[diarize] -> Tai lai diarize_audio.py moi nhat tu Drive/HF roi Runtime -> Restart session.\")\n",
+ "else:\n",
+ " print(\"Bỏ qua cài đặt diarization (ENABLE_DIARIZE = False).\")"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# Tải model GGUF — dùng MODELS_REPO + MODEL_FILES từ cell CẤU HÌNH\n",
+ "from huggingface_hub import hf_hub_download\n",
+ "from pathlib import Path\n",
+ "import os\n",
+ "import shutil\n",
+ "\n",
+ "print(f\"[model] repo: {MODELS_REPO}\")\n",
+ "print(f\"[model] files: {MODEL_FILES}\")\n",
+ "\n",
+ "models_dest = MODELS_DIR\n",
+ "if USE_DRIVE_CACHE:\n",
+ " models_dest = f\"{DRIVE_CACHE}/models\"\n",
+ "Path(models_dest).mkdir(parents=True, exist_ok=True)\n",
+ "\n",
+ "paths = {}\n",
+ "for name in MODEL_FILES:\n",
+ " dest_file = Path(models_dest) / name\n",
+ " if dest_file.is_file():\n",
+ " print(f\"[model] cached: {name}\")\n",
+ " paths[name] = dest_file\n",
+ " continue\n",
+ " print(f\"[model] downloading {MODELS_REPO}/{name} ...\")\n",
+ " p = hf_hub_download(\n",
+ " repo_id=MODELS_REPO,\n",
+ " filename=name,\n",
+ " local_dir=models_dest,\n",
+ " local_dir_use_symlinks=False,\n",
+ " )\n",
+ " paths[name] = Path(p)\n",
+ " print(f\" OK: {p}\")\n",
+ "\n",
+ "# Symlink/copy vào ROOT/models cho translate_srt.py (luôn refresh link hỏng)\n",
+ "Path(MODELS_DIR).mkdir(parents=True, exist_ok=True)\n",
+ "for name, src in paths.items():\n",
+ " src = Path(src).resolve()\n",
+ " dst = Path(MODELS_DIR) / name\n",
+ " if not src.is_file():\n",
+ " raise FileNotFoundError(f\"Model thiếu trên Drive/disk: {src}\")\n",
+ " if dst.is_symlink() or dst.exists():\n",
+ " dst.unlink(missing_ok=True)\n",
+ " try:\n",
+ " os.symlink(src, dst)\n",
+ " except OSError:\n",
+ " if dst.exists():\n",
+ " dst.unlink(missing_ok=True)\n",
+ " shutil.copy2(src, dst)\n",
+ " print(f\"[model] linked {name} -> {dst} ({dst.stat().st_size / 1e9:.1f} GB)\")\n",
+ "\n",
+ "MODEL_PATH = Path(MODELS_DIR) / MODEL_FILES[0]\n",
+ "MMPROJ_PATH = Path(MODELS_DIR) / MODEL_FILES[1]\n",
+ "DRAFT_PATH = Path(MODELS_DIR) / MODEL_FILES[2]\n",
+ "for label, p in [(\"main\", MODEL_PATH), (\"mmproj\", MMPROJ_PATH), (\"draft\", DRAFT_PATH)]:\n",
+ " rp = p.resolve() if p.exists() else p\n",
+ " if not rp.is_file():\n",
+ " raise FileNotFoundError(f\"{label} chưa sẵn sàng: {p}\")\n",
+ "print(f\"\\nModels OK — chạy pipeline được.\")"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# llama-server — restore/build (pin b9553 cho MTP Gemma 4)\n",
+ "import sys\n",
+ "\n",
+ "sys.path.insert(0, ROOT)\n",
+ "from scripts.ensure_llama_colab import ensure_llama_server\n",
+ "\n",
+ "LLAMA_SERVER = str(ensure_llama_server(\n",
+ " llama_dir=LLAMA_DIR,\n",
+ " drive_cache=DRIVE_CACHE if USE_DRIVE_CACHE else None,\n",
+ " allow_build=True,\n",
+ " require_mtp=True,\n",
+ " force_rebuild=FORCE_LLAMA_REBUILD,\n",
+ "))\n",
+ "print(f\"LLAMA_SERVER = {LLAMA_SERVER}\")\n"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# Kiểm tra ffmpeg\n",
+ "!ffmpeg -version | head -1"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# Kiểm tra danh sách JOBS (Drive) hoặc upload 1 phim qua widget\n",
+ "from pathlib import Path\n",
+ "from google.colab import files\n",
+ "import os\n",
+ "\n",
+ "UPLOAD_DIR = \"/content/uploads\"\n",
+ "os.makedirs(UPLOAD_DIR, exist_ok=True)\n",
+ "resolved_jobs = []\n",
+ "\n",
+ "def _list_phim_hint() -> str:\n",
+ " phim = Path(DRIVE_PHIM_DIR)\n",
+ " if not phim.is_dir():\n",
+ " return f\"Chưa có thư mục: {DRIVE_PHIM_DIR}\"\n",
+ " lines = [\"File hiện có trong Gemma/Phim:\"]\n",
+ " for p in sorted(phim.iterdir()):\n",
+ " if p.is_file():\n",
+ " lines.append(f\" - {p.name}\")\n",
+ " if len(lines) == 1:\n",
+ " lines.append(\" (trống — upload video + .srt vào đây)\")\n",
+ " return \"\\n\".join(lines)\n",
+ "\n",
+ "def _find_video(stem: str, preferred: str | None = None) -> Path | None:\n",
+ " if preferred:\n",
+ " p = Path(preferred)\n",
+ " if p.is_file():\n",
+ " return p\n",
+ " for ext in VIDEO_EXTS:\n",
+ " p = Path(DRIVE_PHIM_DIR) / f\"{stem}{ext}\"\n",
+ " if p.is_file():\n",
+ " return p\n",
+ " return None\n",
+ "\n",
+ "if JOBS:\n",
+ " for i, entry in enumerate(JOBS, 1):\n",
+ " stem = entry.get(\"stem\") or Path(entry[\"video\"]).stem\n",
+ " v = _find_video(stem, entry.get(\"video\"))\n",
+ " s = Path(entry[\"srt\"])\n",
+ " if v is None:\n",
+ " raise FileNotFoundError(\n",
+ " f\"Job {i}: không thấy video cho '{stem}'\\n\"\n",
+ " f\"Đã thử: {', '.join(stem + ext for ext in VIDEO_EXTS)}\\n\\n\"\n",
+ " f\"Tên trong PHIM_STEMS phải KHỚP Y HỆT tên file (không đuôi).\\n\\n\"\n",
+ " f\"{_list_phim_hint()}\"\n",
+ " )\n",
+ " if not s.is_file():\n",
+ " s = Path(DRIVE_PHIM_DIR) / f\"{stem}.srt\"\n",
+ " if not s.is_file():\n",
+ " raise FileNotFoundError(\n",
+ " f\"Job {i}: không thấy SRT: {s}\\n\\n\"\n",
+ " f\"Cần file: {stem}.srt cùng thư mục Phim.\\n\\n\"\n",
+ " f\"{_list_phim_hint()}\"\n",
+ " )\n",
+ " out = Path(entry[\"output\"]) if entry.get(\"output\") else v.with_name(v.stem + \".vi.srt\")\n",
+ " resolved_jobs.append({\"video\": v, \"srt\": s, \"output\": out})\n",
+ " print(f\"Job {i}/{len(JOBS)}: {v.name} + {s.name} -> {out.name}\")\n",
+ "elif UPLOAD_WIDGET:\n",
+ " print(\"Upload 1 video (.mp4/.mkv) và 1 SRT (.srt):\")\n",
+ " uploaded = files.upload()\n",
+ " video = srt = None\n",
+ " for name in uploaded:\n",
+ " p = Path(UPLOAD_DIR) / name\n",
+ " p.write_bytes(uploaded[name])\n",
+ " low = name.lower()\n",
+ " if low.endswith((\".mp4\", \".mkv\", \".avi\", \".mov\")):\n",
+ " video = p\n",
+ " elif low.endswith(\".srt\"):\n",
+ " srt = p\n",
+ " if not video or not srt:\n",
+ " raise ValueError(\"Cần 1 file video và 1 file .srt\")\n",
+ " out = video.with_name(video.stem + \".vi.srt\")\n",
+ " resolved_jobs.append({\"video\": video, \"srt\": srt, \"output\": out})\n",
+ " print(f\"Job 1/1: {video.name} -> {out.name}\")\n",
+ "else:\n",
+ " raise ValueError(\"Thêm phim vào JOBS hoặc đặt UPLOAD_WIDGET = True\")\n",
+ "\n",
+ "print(f\"\\nTổng: {len(resolved_jobs)} phim (chạy tuần tự)\")"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# ═══ CHẠY DỊCH SRT (tuần tự từng phim) ═══\n",
+ "import sys\n",
+ "import os\n",
+ "from pathlib import Path\n",
+ "\n",
+ "sys.path.insert(0, ROOT)\n",
+ "from scripts.ensure_llama_colab import ensure_llama_server, llama_bin_valid\n",
+ "\n",
+ "if not llama_bin_valid(Path(LLAMA_SERVER)):\n",
+ " LLAMA_SERVER = str(ensure_llama_server(\n",
+ " llama_dir=LLAMA_DIR,\n",
+ " drive_cache=DRIVE_CACHE if USE_DRIVE_CACHE else None,\n",
+ " allow_build=True,\n",
+ " require_mtp=True,\n",
+ " force_rebuild=FORCE_LLAMA_REBUILD,\n",
+ " ))\n",
+ "\n",
+ "from translate_srt import build_parser, run_pipeline\n",
+ "\n",
+ "def _require_file(path, hint: str) -> Path:\n",
+ " p = Path(path)\n",
+ " if p.is_symlink():\n",
+ " p = p.resolve()\n",
+ " if not p.is_file():\n",
+ " raise FileNotFoundError(f\"{hint}\\n Path: {path}\")\n",
+ " return p\n",
+ "\n",
+ "# Preflight — tránh chạy 2911 cue rồi mới lỗi thiếu model\n",
+ "_model = _require_file(MODEL_PATH, \"Thiếu model GGUF — chạy cell 'Tải model GGUF' (~12GB).\")\n",
+ "_mmproj = _require_file(MMPROJ_PATH, \"Thiếu mmproj — chạy cell 'Tải model GGUF'.\")\n",
+ "_draft = _require_file(DRAFT_PATH, \"Thiếu MTP draft — chạy cell 'Tải model GGUF'.\")\n",
+ "print(f\"[ok] model {_model.name} ({_model.stat().st_size / 1e9:.1f} GB)\")\n",
+ "\n",
+ "if ENABLE_DIARIZE:\n",
+ " try:\n",
+ " from pyannote.audio import Pipeline # noqa: F401\n",
+ " except ImportError as e:\n",
+ " raise RuntimeError(\n",
+ " f\"ENABLE_DIARIZE=True nhưng pyannote import lỗi: {e}\\n\"\n",
+ " \"Chạy lại cell cài Pass 0 hoặc Runtime → Restart session.\"\n",
+ " ) from e\n",
+ "\n",
+ "completed = []\n",
+ "\n",
+ "for i, job in enumerate(resolved_jobs, 1):\n",
+ " v, s, out = job[\"video\"], job[\"srt\"], job[\"output\"]\n",
+ " print(\"\\n\" + \"=\" * 50)\n",
+ " print(f\"=== Phim {i}/{len(resolved_jobs)}: {v.name} ===\")\n",
+ " print(\"=\" * 50)\n",
+ "\n",
+ " argv = [\n",
+ " \"--video\", str(v),\n",
+ " \"--input-srt\", str(s),\n",
+ " \"--output-srt\", str(out),\n",
+ " \"--source-lang\", SOURCE_LANG,\n",
+ " \"--target-lang\", TARGET_LANG,\n",
+ " \"--llama-server\", LLAMA_SERVER,\n",
+ " \"--model\", str(_model.resolve()),\n",
+ " \"--mmproj\", str(_mmproj.resolve()),\n",
+ " \"--model-draft\", str(_draft.resolve()),\n",
+ " \"--limit\", str(LIMIT_CUES),\n",
+ " \"--ngl\", \"999\",\n",
+ " \"--ctx\", str(CTX_SIZE),\n",
+ " \"--scene-frames\", str(SCENE_FRAMES),\n",
+ " \"--scene-max-cues\", str(SCENE_MAX_CUES),\n",
+ " \"--scene-max-gap\", str(SCENE_MAX_GAP),\n",
+ " \"--scene-max-dur\", str(SCENE_MAX_DUR),\n",
+ " ]\n",
+ " if SKIP_CORRECTION:\n",
+ " argv.append(\"--skip-correction\")\n",
+ " if NO_MTP:\n",
+ " argv.append(\"--no-mtp\")\n",
+ " if NO_SCENE_VISION:\n",
+ " argv.append(\"--no-scene-vision\")\n",
+ " if NO_SCENE_IMAGE:\n",
+ " argv.append(\"--no-scene-image\")\n",
+ " if ENABLE_DIARIZE:\n",
+ " argv += [\n",
+ " \"--diarize\",\n",
+ " \"--gender-method\", GENDER_METHOD,\n",
+ " \"--diarize-device\", DIARIZE_DEVICE,\n",
+ " ]\n",
+ " if NUM_SPEAKERS and int(NUM_SPEAKERS) > 0:\n",
+ " argv += [\"--num-speakers\", str(int(NUM_SPEAKERS))]\n",
+ " if HF_TOKEN.strip():\n",
+ " argv += [\"--hf-token\", HF_TOKEN.strip()]\n",
+ " if not DETECT_GENDER:\n",
+ " argv.append(\"--no-gender\")\n",
+ "\n",
+ " args = build_parser().parse_args(argv)\n",
+ " args.log_fn = lambda msg, _i=i: print(f\"[{_i}] {msg}\", flush=True)\n",
+ " run_pipeline(args)\n",
+ " completed.append(out)\n",
+ " print(f\"\\nXong phim {i}: {out}\")\n",
+ "\n",
+ "print(\"\\n\" + \"=\" * 50)\n",
+ "print(f\"HOÀN TẤT {len(completed)} phim:\")\n",
+ "for p in completed:\n",
+ " print(f\" - {p}\")"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# Tải kết quả về máy (tuỳ chọn) + ngắt kết nối\n",
+ "from pathlib import Path\n",
+ "import time\n",
+ "\n",
+ "if DOWNLOAD_RESULTS:\n",
+ " from google.colab import files\n",
+ " for out in completed:\n",
+ " out = Path(out)\n",
+ " for p in [\n",
+ " out,\n",
+ " out.with_suffix(\".corrected\" + out.suffix),\n",
+ " out.with_suffix(out.suffix + \".report.json\"),\n",
+ " ]:\n",
+ " if p.is_file():\n",
+ " print(f\"Download: {p.name}\")\n",
+ " files.download(str(p))\n",
+ "else:\n",
+ " print(\"Bỏ qua download — file SRT đã nằm trên Drive (xem đường dẫn output trong JOBS).\")\n",
+ "\n",
+ "if AUTO_UNMOUNT_DRIVE and USE_DRIVE_CACHE:\n",
+ " from google.colab import drive\n",
+ " drive.flush_and_unmount()\n",
+ " print(\"Đã gỡ mount Google Drive.\")\n",
+ "\n",
+ "if AUTO_DISCONNECT_RUNTIME:\n",
+ " delay = max(0, int(DISCONNECT_DELAY_SEC))\n",
+ " if delay:\n",
+ " print(f\"Ngắt Colab runtime sau {delay}s...\")\n",
+ " time.sleep(delay)\n",
+ " from google.colab import runtime\n",
+ " runtime.unassign()"
+ ],
+ "execution_count": null,
+ "outputs": []
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "L4",
+ "name": "GemmaSRT_Colab.ipynb",
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
\ No newline at end of file
diff --git a/hf-upload/config.yaml b/hf-upload/config.yaml
new file mode 100644
index 0000000000000000000000000000000000000000..ba276264db1881904ff2d9402c701f1ad736c87b
--- /dev/null
+++ b/hf-upload/config.yaml
@@ -0,0 +1,40 @@
+# Cấu hình mặc định cho Gemma SRT Translate (Colab + Hugging Face)
+# Thay YOUR_USERNAME trước khi upload repo.
+
+repo:
+ id: "STBack23/gemma-srt-translate"
+ type: "model" # hoặc "space" nếu deploy Space sau này
+
+models:
+ source_repo: "unsloth/gemma-4-12B-it-qat-GGUF"
+ files:
+ - "gemma-4-12B-it-qat-UD-Q4_K_XL.gguf"
+ - "mmproj-F16.gguf"
+ - "mtp-gemma-4-12B-it.gguf"
+
+llama_cpp:
+ git_url: "https://github.com/ggml-org/llama.cpp.git"
+ cuda_arch: "89" # L4, RTX 4090 (Ada)
+ build_target: "llama-server"
+
+translate:
+ source_lang: "auto"
+ target_lang: "Vietnamese"
+ chars_per_sec: 22.0
+ max_line_chars: 52
+ max_lines: 2
+ shorten: false
+ scene_max_gap: 1.5
+ scene_max_cues: 4
+ scene_max_dur: 20.0
+ ctx: 8192
+ ngl: 999
+ temp: 0.7
+ use_mtp: true
+
+paths:
+ local_models: "models"
+ colab_root: "/content/gemma-srt"
+ drive_root: "/content/drive/MyDrive/Gemma"
+ drive_cache: "/content/drive/MyDrive/Gemma/Cache"
+ drive_phim: "/content/drive/MyDrive/Gemma/Phim"
diff --git a/hf-upload/diarize_audio.py b/hf-upload/diarize_audio.py
new file mode 100644
index 0000000000000000000000000000000000000000..5a93dc705397656a10a360145b181b22f9c3e2a6
--- /dev/null
+++ b/hf-upload/diarize_audio.py
@@ -0,0 +1,1130 @@
+"""Speaker diarization + gender estimation for subtitle cues.
+
+This module answers "who spoke when" using pyannote.audio and attaches a stable
+speaker label (and a coarse male/female/unknown gender guess) to every SRT cue.
+Gemma then uses those labels to keep Vietnamese forms of address / pronouns
+(ta, hắn, y, muội muội, tỷ tỷ, phu quân, ...) consistent for each character.
+
+Pipeline:
+ extract_audio (ffmpeg -> 16k mono wav) ->
+ run_diarization (pyannote community-1) -> list[SpeakerTurn] ->
+ estimate_speaker_genders (median F0 per speaker, optional) ->
+ assign_speakers_to_cues (max temporal overlap)
+
+Heavy dependencies (pyannote.audio, torch, torchaudio) are imported lazily so
+that the rest of the project keeps working even when they are not installed.
+The Hugging Face access token is required by pyannote and is read, in order,
+from the explicit argument, then the ``HF_TOKEN`` / ``HUGGINGFACE_TOKEN`` env
+vars.
+"""
+
+from __future__ import annotations
+
+import os
+import subprocess
+import sys
+import tempfile
+import time
+from dataclasses import dataclass
+from pathlib import Path
+from typing import TYPE_CHECKING, Any, Callable
+
+if TYPE_CHECKING: # pragma: no cover - typing only
+ from translate_srt import Cue
+
+# Preferred open-source pipeline (pyannote.audio 4.x). Falls back to the legacy
+# 3.1 pipeline if the newer one is unavailable for the installed version.
+DEFAULT_PIPELINE = "pyannote/speaker-diarization-community-1"
+FALLBACK_PIPELINE = "pyannote/speaker-diarization-3.1"
+
+def _load_default_token() -> str:
+ """Local convenience token, read from a sibling file that is NOT uploaded.
+
+ Put your token in ``hf_token.local`` (or ``hf_token.txt``) next to this file
+ to have local runs pick it up automatically. The file is git-ignored and is
+ never copied by ``prepare_upload.ps1`` so the public code stays secret-free.
+ """
+ for name in ("hf_token.local", "hf_token.txt"):
+ p = Path(__file__).with_name(name)
+ try:
+ if p.is_file():
+ tok = p.read_text(encoding="utf-8").strip()
+ if tok:
+ return tok
+ except OSError:
+ pass
+ return ""
+
+
+# Built-in fallback Hugging Face token (used when no arg/env token is supplied).
+# Loaded from a local-only file so the public source never embeds a secret.
+DEFAULT_HF_TOKEN = _load_default_token()
+
+_SUBPROCESS_FLAGS = subprocess.CREATE_NO_WINDOW if sys.platform == "win32" else 0
+
+# Coarse pitch thresholds (Hz) for median fundamental frequency per speaker.
+# A deliberate "unknown" band avoids confidently mislabelling ambiguous voices.
+# Male speech F0 is typically ~85-180 Hz and female ~165-255 Hz, so the bands
+# overlap. We label only confident cases and leave the overlap as "unknown" so
+# the translator infers gender from dialogue/context instead of guessing wrong.
+_MALE_F0_MAX = 150.0
+_FEMALE_F0_MIN = 195.0
+_F0_MIN_HZ = 65.0
+_F0_MAX_HZ = 400.0
+
+# Dedicated speech model that predicts age (0..100) and gender (female/male/child).
+DEFAULT_AGE_GENDER_MODEL = "audeering/wav2vec2-large-robust-24-ft-age-gender"
+# The model's gender head order is [female, male, child].
+_AGE_GENDER_LABELS = ("female", "male", "child")
+# Minimum softmax confidence to commit to a gender label; below this we report
+# "unknown" so the translator infers gender from dialogue instead of trusting a
+# coin-flip. Mirrors the deliberate "unknown" band used by the pitch heuristic.
+_MODEL_GENDER_MIN_CONF = 0.55
+# The model often mislabels high-pitched adult women as "child"; only keep the
+# "child" gender when the predicted age is genuinely childlike.
+_CHILD_MAX_AGE = 16.0
+
+# Speaker-embedding model used to merge speakers that diarization over-split.
+# speechbrain ECAPA is not gated (no extra HF accept needed) and ships as a
+# pyannote.audio dependency.
+DEFAULT_EMBEDDING_MODEL = "speechbrain/spkrec-ecapa-voxceleb"
+# Conservative cosine-similarity threshold: only merge two labels when we are
+# confident they are the SAME voice that pyannote wrongly split apart.
+_MERGE_SIM_THRESHOLD = 0.70
+
+# Cached (model, processor, device) so we only load the ~1 GB model once.
+_AGE_GENDER_CACHE: dict[str, Any] = {}
+
+
+def age_to_group(age_years: float | None, gender: str = "") -> str:
+ """Bucket an age (years) into a coarse life-stage label.
+
+ Returns one of "child", "teen", "young_adult", "middle_aged", "elderly" or
+ "" when age is unknown.
+ """
+ if gender == "child":
+ return "child"
+ if age_years is None:
+ return ""
+ if age_years < 13:
+ return "child"
+ if age_years < 20:
+ return "teen"
+ if age_years < 40:
+ return "young_adult"
+ if age_years < 60:
+ return "middle_aged"
+ return "elderly"
+
+
+LogFn = Callable[[str], None]
+
+
+def _noop(_msg: str) -> None:
+ pass
+
+
+def _fmt_dur(seconds: float) -> str:
+ """Human-readable duration, e.g. '1h 23m' or '4m 12s'."""
+ if seconds < 0:
+ seconds = 0.0
+ s = int(round(seconds))
+ if s >= 3600:
+ h, rem = divmod(s, 3600)
+ m, sec = divmod(rem, 60)
+ return f"{h}h {m}m" if sec < 30 else f"{h}h {m}m {sec}s"
+ if s >= 60:
+ m, sec = divmod(s, 60)
+ return f"{m}m {sec}s"
+ return f"{s}s"
+
+
+def _log_step(log: LogFn, step: int, total: int, msg: str) -> None:
+ log(f"[diarize] [{step}/{total}] {msg}")
+
+
+def _audio_duration_sec(waveform, sample_rate: int) -> float:
+ try:
+ n = int(waveform.shape[-1])
+ return n / float(sample_rate) if sample_rate else 0.0
+ except Exception: # noqa: BLE001
+ return 0.0
+
+
+@dataclass
+class SpeakerTurn:
+ start: float
+ end: float
+ speaker: str
+
+ @property
+ def duration(self) -> float:
+ return max(0.0, self.end - self.start)
+
+
+# --------------------------------------------------------------------------- #
+# Token handling
+# --------------------------------------------------------------------------- #
+def resolve_hf_token(token: str | None) -> str | None:
+ if token and token.strip():
+ return token.strip()
+ for env in ("HF_TOKEN", "HUGGINGFACE_TOKEN", "HUGGING_FACE_HUB_TOKEN"):
+ val = os.environ.get(env)
+ if val and val.strip():
+ return val.strip()
+ if DEFAULT_HF_TOKEN and DEFAULT_HF_TOKEN.strip():
+ return DEFAULT_HF_TOKEN.strip()
+ return None
+
+
+# --------------------------------------------------------------------------- #
+# Audio extraction
+# --------------------------------------------------------------------------- #
+def extract_audio(
+ video: Path,
+ out_wav: Path,
+ sample_rate: int = 16000,
+ max_seconds: float | None = None,
+) -> Path:
+ """Extract a 16 kHz mono PCM wav from *video* (what pyannote expects).
+
+ When *max_seconds* is given, only that many seconds from the start are
+ extracted (so diarization scales with the number of cues being processed).
+ """
+ cmd = [
+ "ffmpeg",
+ "-hide_banner",
+ "-loglevel",
+ "error",
+ "-i",
+ str(video),
+ "-vn",
+ "-ac",
+ "1",
+ "-ar",
+ str(sample_rate),
+ ]
+ if max_seconds and max_seconds > 0:
+ cmd += ["-t", f"{max_seconds:.3f}"]
+ cmd += [
+ "-c:a",
+ "pcm_s16le",
+ "-y",
+ str(out_wav),
+ ]
+ proc = subprocess.run(
+ cmd, capture_output=True, timeout=600, creationflags=_SUBPROCESS_FLAGS
+ )
+ if proc.returncode != 0:
+ raise RuntimeError(
+ "ffmpeg audio extraction failed: "
+ + proc.stderr.decode("utf-8", "replace").strip()
+ )
+ if not out_wav.exists():
+ raise RuntimeError(f"Audio extraction produced no file at {out_wav}")
+ return out_wav
+
+
+# --------------------------------------------------------------------------- #
+# Diarization
+# --------------------------------------------------------------------------- #
+def _pick_device(preferred: str = "auto") -> str:
+ if preferred and preferred != "auto":
+ return preferred
+ try:
+ import torch # noqa: PLC0415
+
+ return "cuda" if torch.cuda.is_available() else "cpu"
+ except Exception: # noqa: BLE001
+ return "cpu"
+
+
+def _patch_speechbrain_lazy(log: LogFn = _noop) -> None:
+ """Work around a speechbrain LazyModule bug that breaks pyannote on Windows.
+
+ ``speechbrain.utils.importutils.LazyModule.ensure_module`` tries to skip the
+ lazy import when the caller is ``inspect.py`` (e.g. ``inspect.getmodule``
+ probing ``__file__``), but it only matches the POSIX ``/inspect.py`` path. On
+ Windows the path uses ``\\`` so the guard misses, the lazy module (e.g.
+ ``speechbrain.integrations.k2_fsa``) is force-imported, and it raises
+ ``ImportError`` — which ``hasattr`` propagates, crashing pipeline loading.
+
+ We re-implement ``ensure_module`` with an OS-agnostic ``inspect.py`` check.
+ """
+ try:
+ import importlib as _importlib # noqa: PLC0415
+ import inspect as _inspect # noqa: PLC0415
+ import os as _os # noqa: PLC0415
+ import sys as _sys # noqa: PLC0415
+
+ from speechbrain.utils import importutils as iu # noqa: PLC0415
+ except Exception: # noqa: BLE001 - speechbrain not used by this pipeline
+ return
+
+ if getattr(iu.LazyModule, "_gemma_patched", False):
+ return
+
+ def ensure_module(self, stacklevel: int):
+ importer_frame = None
+ try:
+ importer_frame = _inspect.getframeinfo(_sys._getframe(stacklevel + 1))
+ except (AttributeError, ValueError):
+ importer_frame = None
+ if (
+ importer_frame is not None
+ and _os.path.basename(importer_frame.filename) == "inspect.py"
+ ):
+ raise AttributeError()
+ if self.lazy_module is None:
+ try:
+ if self.package is None:
+ self.lazy_module = _importlib.import_module(self.target)
+ else:
+ self.lazy_module = _importlib.import_module(
+ f".{self.target}", self.package
+ )
+ except Exception as e: # noqa: BLE001
+ raise ImportError(f"Lazy import of {self!r} failed") from e
+ return self.lazy_module
+
+ iu.LazyModule.ensure_module = ensure_module
+ iu.LazyModule._gemma_patched = True
+ log("[diarize] applied speechbrain LazyModule Windows compatibility patch.")
+
+
+def _load_pipeline(hf_token: str, device: str, log: LogFn):
+ try:
+ from pyannote.audio import Pipeline # noqa: PLC0415
+ except ImportError as exc: # pragma: no cover - env dependent
+ raise RuntimeError(
+ "pyannote.audio is not installed. Install diarization extras:\n"
+ " pip install -r requirements-diarize.txt"
+ ) from exc
+
+ import torch # noqa: PLC0415
+
+ _patch_speechbrain_lazy(log)
+
+ last_err: Exception | None = None
+ for model_id in (DEFAULT_PIPELINE, FALLBACK_PIPELINE):
+ try:
+ log(f"[diarize] loading pipeline {model_id} ...")
+ pipeline = Pipeline.from_pretrained(model_id, token=hf_token)
+ if pipeline is None:
+ raise RuntimeError(
+ "Pipeline.from_pretrained returned None — the HF token is "
+ "likely invalid or the model conditions were not accepted at "
+ f"https://hf.co/{model_id}"
+ )
+ try:
+ pipeline.to(torch.device(device))
+ except Exception: # noqa: BLE001 - keep CPU pipeline if .to fails
+ pass
+ log(f"[diarize] pipeline ready: {model_id} on {device}.")
+ return pipeline
+ except Exception as exc: # noqa: BLE001
+ last_err = exc
+ log(f"[diarize] could not load {model_id}: {exc}")
+ raise RuntimeError(
+ "Failed to load any diarization pipeline. Most common causes:\n"
+ " 1) HF token's account has not accepted ALL gated model conditions (Agree):\n"
+ f" https://hf.co/{DEFAULT_PIPELINE}\n"
+ f" https://hf.co/{FALLBACK_PIPELINE}\n"
+ " https://hf.co/pyannote/segmentation-3.0\n"
+ " 2) On Colab: pytorch-lightning version mismatch (community-1).\n"
+ f"Last error: {last_err}"
+ )
+
+
+def load_wav(wav: Path):
+ """Load a wav into an in-memory (channel, time) float32 tensor + sample rate.
+
+ Uses ``soundfile`` instead of ``torchaudio.load`` / torchcodec so audio
+ decoding does not depend on a working torchcodec/ffmpeg-DLL setup.
+ """
+ import soundfile as sf # noqa: PLC0415
+ import torch # noqa: PLC0415
+
+ data, sr = sf.read(str(wav), dtype="float32", always_2d=True) # (time, channels)
+ waveform = torch.from_numpy(data.T).contiguous() # (channels, time)
+ if waveform.shape[0] > 1:
+ waveform = waveform.mean(dim=0, keepdim=True)
+ return waveform, int(sr)
+
+
+def _as_annotation(diarization):
+ """Return a pyannote Annotation (with ``itertracks``) from a pipeline result.
+
+ pyannote.audio 4.x (community-1) returns a ``DiarizeOutput`` exposing
+ ``speaker_diarization`` and ``exclusive_speaker_diarization``; the legacy 3.1
+ pipeline returns an ``Annotation`` directly. Exclusive diarization assigns at
+ most one speaker per instant, which is ideal for tagging subtitle cues.
+ """
+ if hasattr(diarization, "itertracks"):
+ return diarization
+ for attr in ("exclusive_speaker_diarization", "speaker_diarization"):
+ candidate = getattr(diarization, attr, None)
+ if candidate is not None and hasattr(candidate, "itertracks"):
+ return candidate
+ raise RuntimeError(
+ f"Unexpected diarization output type without itertracks: {type(diarization)!r}"
+ )
+
+
+def run_diarization(
+ waveform,
+ sample_rate: int,
+ hf_token: str,
+ num_speakers: int | None = None,
+ min_speakers: int | None = None,
+ max_speakers: int | None = None,
+ device: str = "auto",
+ log: LogFn = _noop,
+) -> list[SpeakerTurn]:
+ """Run pyannote diarization and return merged speaker turns sorted by time."""
+ dev = _pick_device(device)
+ pipeline = _load_pipeline(hf_token, dev, log)
+
+ kwargs: dict[str, int] = {}
+ if num_speakers and num_speakers > 0:
+ kwargs["num_speakers"] = num_speakers
+ else:
+ if min_speakers and min_speakers > 0:
+ kwargs["min_speakers"] = min_speakers
+ if max_speakers and max_speakers > 0:
+ kwargs["max_speakers"] = max_speakers
+
+ dur = _audio_duration_sec(waveform, sample_rate)
+ kw_hint = ", ".join(f"{k}={v}" for k, v in kwargs.items()) or "auto-detect speakers"
+ log(
+ f"[diarize] running pyannote on {_fmt_dur(dur)} audio "
+ f"({dev}, {kw_hint}) — có thể mất vài phút đến ~1h với phim dài..."
+ )
+ t0 = time.time()
+ audio = {"waveform": waveform, "sample_rate": sample_rate}
+ diarization = pipeline(audio, **kwargs)
+ log(f"[diarize] pyannote inference done in {_fmt_dur(time.time() - t0)}.")
+ annotation = _as_annotation(diarization)
+
+ turns: list[SpeakerTurn] = []
+ for segment, _track, speaker in annotation.itertracks(yield_label=True):
+ turns.append(SpeakerTurn(float(segment.start), float(segment.end), str(speaker)))
+ turns.sort(key=lambda t: (t.start, t.end))
+ n_spk = len({t.speaker for t in turns})
+ log(f"[diarize] found {n_spk} speaker(s) across {len(turns)} turn(s).")
+ return turns
+
+
+# --------------------------------------------------------------------------- #
+# Gender estimation (coarse, pitch based)
+# --------------------------------------------------------------------------- #
+def estimate_speaker_genders(
+ waveform,
+ sr: int,
+ turns: list[SpeakerTurn],
+ *,
+ max_seconds_per_speaker: float = 30.0,
+ log: LogFn = _noop,
+) -> dict[str, str]:
+ """Estimate male/female/unknown per speaker from median fundamental frequency.
+
+ This is a lightweight heuristic (no extra model download) meant only as a
+ hint for the translator. Returns a mapping ``{speaker_label: gender}`` where
+ gender is one of ``"male"``, ``"female"`` or ``"unknown"``. *waveform* is a
+ ``(channel, time)`` float32 tensor as returned by :func:`load_wav`.
+ """
+ try:
+ import torch # noqa: PLC0415
+ import torchaudio.functional as AF # noqa: PLC0415
+ except Exception as exc: # noqa: BLE001
+ log(f"[diarize] gender estimation skipped (torch/torchaudio missing): {exc}")
+ return {}
+
+ if waveform.ndim == 1:
+ waveform = waveform.unsqueeze(0)
+ elif waveform.shape[0] > 1:
+ waveform = waveform.mean(dim=0, keepdim=True)
+
+ by_speaker: dict[str, list[SpeakerTurn]] = {}
+ for t in turns:
+ by_speaker.setdefault(t.speaker, []).append(t)
+
+ genders: dict[str, str] = {}
+ for speaker, spk_turns in by_speaker.items():
+ spk_turns = sorted(spk_turns, key=lambda t: t.duration, reverse=True)
+ chunks: list["torch.Tensor"] = []
+ used = 0.0
+ for t in spk_turns:
+ if used >= max_seconds_per_speaker:
+ break
+ s = max(0, int(t.start * sr))
+ e = min(waveform.shape[-1], int(t.end * sr))
+ if e - s < int(0.20 * sr):
+ continue
+ chunks.append(waveform[..., s:e])
+ used += (e - s) / sr
+ if not chunks:
+ genders[speaker] = "unknown"
+ continue
+ audio = torch.cat(chunks, dim=-1)
+ try:
+ pitch = AF.detect_pitch_frequency(audio, sr)
+ except Exception as exc: # noqa: BLE001
+ log(f"[diarize] pitch detection failed for {speaker}: {exc}")
+ genders[speaker] = "unknown"
+ continue
+ flat = pitch.flatten()
+ voiced = flat[(flat >= _F0_MIN_HZ) & (flat <= _F0_MAX_HZ)]
+ if voiced.numel() < 5:
+ genders[speaker] = "unknown"
+ continue
+ median_f0 = float(voiced.median().item())
+ if median_f0 <= _MALE_F0_MAX:
+ gender = "male"
+ elif median_f0 >= _FEMALE_F0_MIN:
+ gender = "female"
+ else:
+ gender = "unknown"
+ genders[speaker] = gender
+ log(f"[diarize] {speaker}: median F0 ~{median_f0:.0f} Hz -> {gender}")
+ return genders
+
+
+# --------------------------------------------------------------------------- #
+# Age + gender estimation (dedicated wav2vec2 model)
+# --------------------------------------------------------------------------- #
+def _weighted_median(values: list[float], weights: list[float]) -> float:
+ """Weighted median — robust to a few outlier windows (e.g. noisy segments).
+
+ Falls back to the plain mean when weights are missing/zero.
+ """
+ pairs = [(v, w) for v, w in zip(values, weights) if w and w > 0]
+ if not pairs:
+ return sum(values) / len(values) if values else 0.0
+ pairs.sort(key=lambda p: p[0])
+ total = sum(w for _, w in pairs)
+ half = total / 2.0
+ acc = 0.0
+ for v, w in pairs:
+ acc += w
+ if acc >= half:
+ return v
+ return pairs[-1][0]
+
+
+def _collect_speaker_windows(
+ waveform,
+ sr,
+ turns,
+ *,
+ max_seconds_per_speaker: float = 45.0,
+ window_sec: float = 8.0,
+ min_sec: float = 0.6,
+ max_windows: int = 12,
+):
+ """Group turns per speaker and split them into short coherent windows.
+
+ Returns ``{speaker: [(1-D float32 torch.Tensor, duration_sec), ...]}``.
+
+ Unlike concatenating one long 30 s clip and mean-pooling it once (which
+ biases the predicted age toward the dataset mean and lets a couple of noisy
+ seconds flip the gender), we keep each window separate so the caller can
+ predict per window and aggregate robustly (weighted median age, weighted
+ mean gender probability). Longest turns are used first.
+ """
+ import torch # noqa: PLC0415
+
+ if waveform.ndim == 1:
+ waveform = waveform.unsqueeze(0)
+ elif waveform.shape[0] > 1:
+ waveform = waveform.mean(dim=0, keepdim=True)
+
+ by_speaker: dict[str, list[SpeakerTurn]] = {}
+ for t in turns:
+ by_speaker.setdefault(t.speaker, []).append(t)
+
+ win = max(1, int(window_sec * sr))
+ min_len = max(1, int(min_sec * sr))
+ total_samples = waveform.shape[-1]
+
+ out: dict[str, list] = {}
+ for speaker, spk_turns in by_speaker.items():
+ spk_turns = sorted(spk_turns, key=lambda t: t.duration, reverse=True)
+ windows: list = []
+ used = 0.0
+ for t in spk_turns:
+ if used >= max_seconds_per_speaker or len(windows) >= max_windows:
+ break
+ s = max(0, int(t.start * sr))
+ e = min(total_samples, int(t.end * sr))
+ pos = s
+ while pos + min_len <= e:
+ if used >= max_seconds_per_speaker or len(windows) >= max_windows:
+ break
+ w_end = min(e, pos + win)
+ if w_end - pos < min_len:
+ break
+ seg = waveform[..., pos:w_end].flatten()
+ dur = (w_end - pos) / sr
+ windows.append((seg, dur))
+ used += dur
+ pos = w_end
+ if windows:
+ out[speaker] = windows
+ return out
+
+
+def _load_age_gender_model(model_name: str, device: str, log: LogFn):
+ """Load (and cache) the audeering age/gender wav2vec2 model + processor."""
+ cache_key = f"{model_name}@{device}"
+ if cache_key in _AGE_GENDER_CACHE:
+ return _AGE_GENDER_CACHE[cache_key]
+
+ import torch # noqa: PLC0415
+ import torch.nn as nn # noqa: PLC0415
+ from transformers import Wav2Vec2Processor # noqa: PLC0415
+ from transformers.models.wav2vec2.modeling_wav2vec2 import ( # noqa: PLC0415
+ Wav2Vec2Model,
+ Wav2Vec2PreTrainedModel,
+ )
+
+ class ModelHead(nn.Module):
+ def __init__(self, config, num_labels):
+ super().__init__()
+ self.dense = nn.Linear(config.hidden_size, config.hidden_size)
+ self.dropout = nn.Dropout(config.final_dropout)
+ self.out_proj = nn.Linear(config.hidden_size, num_labels)
+
+ def forward(self, features):
+ x = self.dropout(features)
+ x = self.dense(x)
+ x = torch.tanh(x)
+ x = self.dropout(x)
+ return self.out_proj(x)
+
+ class AgeGenderModel(Wav2Vec2PreTrainedModel):
+ # This model has no tied weights, so the mapping is always empty.
+ _tied_weights_keys = {}
+
+ def __init__(self, config):
+ super().__init__(config)
+ self.wav2vec2 = Wav2Vec2Model(config)
+ self.age = ModelHead(config, 1)
+ self.gender = ModelHead(config, 3)
+ # transformers >=5 builds `all_tied_weights_keys` in post_init();
+ # older releases only have init_weights(). Support both.
+ if hasattr(self, "post_init"):
+ self.post_init()
+ else:
+ self.init_weights()
+ # `from_pretrained` calls `all_tied_weights_keys.keys()` while
+ # loading; guarantee it is a dict so the load never crashes.
+ if not isinstance(getattr(self, "all_tied_weights_keys", None), dict):
+ self.all_tied_weights_keys = {}
+
+ def forward(self, input_values):
+ outputs = self.wav2vec2(input_values)
+ hidden_states = outputs[0]
+ hidden_states = torch.mean(hidden_states, dim=1)
+ logits_age = self.age(hidden_states)
+ logits_gender = torch.softmax(self.gender(hidden_states), dim=1)
+ return hidden_states, logits_age, logits_gender
+
+ log(f"[diarize] loading age/gender model {model_name} (first run downloads ~1 GB)...")
+ processor = Wav2Vec2Processor.from_pretrained(model_name)
+ model = AgeGenderModel.from_pretrained(model_name)
+ tied = getattr(model, "all_tied_weights_keys", None)
+ if not isinstance(tied, dict):
+ model.all_tied_weights_keys = {}
+ model = model.to(torch.device(device)).eval()
+ _AGE_GENDER_CACHE[cache_key] = (model, processor, device)
+ log(f"[diarize] age/gender model ready on {device}.")
+ return _AGE_GENDER_CACHE[cache_key]
+
+
+def estimate_speaker_demographics_model(
+ waveform,
+ sr: int,
+ turns: list[SpeakerTurn],
+ *,
+ model_name: str = DEFAULT_AGE_GENDER_MODEL,
+ device: str = "auto",
+ max_seconds_per_speaker: float = 30.0,
+ log: LogFn = _noop,
+) -> dict[str, dict[str, Any]]:
+ """Predict gender (female/male/child) + age per speaker with a dedicated model.
+
+ Returns ``{speaker: {"gender": str, "age": float|None, "age_group": str}}``.
+ """
+ import torch # noqa: PLC0415
+
+ dev = _pick_device(device)
+ model, processor, dev = _load_age_gender_model(model_name, dev, log)
+ windows_by_speaker = _collect_speaker_windows(
+ waveform, sr, turns, max_seconds_per_speaker=max_seconds_per_speaker
+ )
+ n_spk = len(windows_by_speaker)
+ log(f"[diarize] age/gender: predicting {n_spk} speaker(s) (per-window)...")
+
+ out: dict[str, dict[str, Any]] = {}
+ for i, (speaker, windows) in enumerate(windows_by_speaker.items(), start=1):
+ ages: list[float] = []
+ weights: list[float] = []
+ gender_prob_sum = torch.zeros(len(_AGE_GENDER_LABELS))
+ for seg, dur in windows:
+ signal = seg.detach().cpu().numpy().astype("float32")
+ inputs = processor(signal, sampling_rate=sr)
+ values = torch.from_numpy(
+ inputs["input_values"][0].reshape(1, -1)
+ ).to(torch.device(dev))
+ with torch.no_grad():
+ _hidden, logits_age, logits_gender = model(values)
+ ages.append(float(logits_age[0].item()) * 100.0)
+ weights.append(dur)
+ gender_prob_sum += logits_gender[0].detach().cpu().float() * dur
+
+ total_w = sum(weights)
+ if total_w <= 0:
+ continue
+ mean_gender = gender_prob_sum / total_w
+ probs = {lab: float(mean_gender[k].item()) for k, lab in enumerate(_AGE_GENDER_LABELS)}
+ # Weighted median age is robust to a few outlier windows; mean-pooling
+ # one long clip instead pulled every speaker toward the dataset mean.
+ age_years = _weighted_median(ages, weights)
+ raw_gender = max(probs, key=probs.get)
+ gender = raw_gender
+ # Fix the common "adult woman -> child" mislabel: drop child when the
+ # predicted age is clearly adult, choosing the better of female/male.
+ if gender == "child" and age_years >= _CHILD_MAX_AGE:
+ gender = "female" if probs["female"] >= probs["male"] else "male"
+ # Don't commit to a low-confidence guess; "unknown" lets the translator
+ # infer gender from dialogue/context instead.
+ if gender != "child" and probs[gender] < _MODEL_GENDER_MIN_CONF:
+ gender = "unknown"
+ age_group = age_to_group(age_years, gender)
+ out[speaker] = {
+ "gender": gender,
+ "age": round(age_years, 1),
+ "age_group": age_group,
+ }
+ demoted = "" if gender == raw_gender else f" (raw={raw_gender})"
+ log(
+ f"[diarize] age/gender: [{i}/{n_spk}] {speaker}: ~{age_years:.0f}y, "
+ f"gender={gender}{demoted} (conf {probs[raw_gender]:.2f}) "
+ f"({age_group or 'n/a'}) from {len(windows)} window(s)/{total_w:.0f}s"
+ )
+ return out
+
+
+def estimate_speaker_demographics(
+ waveform,
+ sr: int,
+ turns: list[SpeakerTurn],
+ *,
+ method: str = "auto",
+ model_name: str = DEFAULT_AGE_GENDER_MODEL,
+ device: str = "auto",
+ log: LogFn = _noop,
+) -> dict[str, dict[str, Any]]:
+ """Estimate per-speaker demographics, dispatching by *method*.
+
+ method:
+ "model" -> dedicated wav2vec2 age/gender model (accurate, age + gender).
+ "pitch" -> lightweight F0 heuristic (gender only, no age).
+ "auto" -> try the model, fall back to pitch on any failure.
+ Returns ``{speaker: {"gender","age","age_group"}}``.
+ """
+ if method in ("model", "auto"):
+ try:
+ log(f"[diarize] age/gender: loading model ({model_name})...")
+ return estimate_speaker_demographics_model(
+ waveform, sr, turns, model_name=model_name, device=device, log=log
+ )
+ except Exception as exc: # noqa: BLE001
+ log(f"[diarize] age/gender model failed, using pitch heuristic: {exc}")
+
+ log("[diarize] age/gender: pitch heuristic (F0) per speaker...")
+ genders = estimate_speaker_genders(waveform, sr, turns, log=log)
+ return {
+ spk: {"gender": g, "age": None, "age_group": ""} for spk, g in genders.items()
+ }
+
+
+# --------------------------------------------------------------------------- #
+# Speaker merge by voice embedding (fix over-segmentation)
+# --------------------------------------------------------------------------- #
+def _patch_symlink_fallback(log: LogFn = _noop) -> None:
+ """Make os.symlink fall back to copying on Windows without admin rights.
+
+ speechbrain (used by the ECAPA embedding model) symlinks files from the HF
+ cache into its save dir. On Windows this raises ``WinError 1314`` unless
+ Developer Mode or admin is enabled, which aborted the speaker merge. Copying
+ instead is harmless for these small model files.
+ """
+ if getattr(os, "_gemma_symlink_patched", False):
+ return
+ import shutil # noqa: PLC0415
+
+ _orig_symlink = os.symlink
+
+ def _safe_symlink(src, dst, *args, **kwargs): # type: ignore[no-untyped-def]
+ try:
+ return _orig_symlink(src, dst, *args, **kwargs)
+ except OSError:
+ src_s, dst_s = os.fspath(src), os.fspath(dst)
+ if os.path.isdir(src_s):
+ shutil.copytree(src_s, dst_s, dirs_exist_ok=True)
+ else:
+ os.makedirs(os.path.dirname(dst_s) or ".", exist_ok=True)
+ shutil.copyfile(src_s, dst_s)
+ return None
+
+ os.symlink = _safe_symlink # type: ignore[assignment]
+ os._gemma_symlink_patched = True # type: ignore[attr-defined]
+ log("[diarize] applied Windows symlink-copy fallback for embedding model.")
+
+
+def _speaker_embeddings(
+ waveform,
+ sr: int,
+ turns: list[SpeakerTurn],
+ *,
+ model_name: str,
+ device: str,
+ max_seconds_per_speaker: float = 30.0,
+ log: LogFn = _noop,
+) -> dict[str, Any]:
+ """Compute one mean, L2-normalized voice embedding per diarization label."""
+ import numpy as np # noqa: PLC0415
+ import torch # noqa: PLC0415
+
+ os.environ.setdefault("HF_HUB_DISABLE_SYMLINKS_WARNING", "1")
+ _patch_symlink_fallback(log)
+
+ from pyannote.audio.pipelines.speaker_verification import ( # noqa: PLC0415
+ PretrainedSpeakerEmbedding,
+ )
+
+ dev = _pick_device(device)
+ # A persistent cache dir avoids speechbrain writing into a literal "None"
+ # folder (its default when cache_dir is None) and re-downloading each run.
+ cache_dir = Path(tempfile.gettempdir()) / "gemma_speechbrain"
+ cache_dir.mkdir(parents=True, exist_ok=True)
+ try:
+ emb_model = PretrainedSpeakerEmbedding(
+ model_name, device=torch.device(dev), cache_dir=str(cache_dir)
+ )
+ except TypeError:
+ # Older pyannote signatures may not accept cache_dir.
+ emb_model = PretrainedSpeakerEmbedding(model_name, device=torch.device(dev))
+ windows_by_speaker = _collect_speaker_windows(
+ waveform, sr, turns, max_seconds_per_speaker=max_seconds_per_speaker
+ )
+ out: dict[str, Any] = {}
+ for speaker, windows in windows_by_speaker.items():
+ vecs: list[Any] = []
+ for seg, _dur in windows:
+ wav = seg.detach().float().reshape(1, 1, -1)
+ try:
+ emb = emb_model(wav)
+ except Exception: # noqa: BLE001 - skip too-short/odd windows
+ continue
+ v = np.asarray(emb, dtype="float64").reshape(-1)
+ norm = float(np.linalg.norm(v))
+ if norm > 0 and np.isfinite(norm):
+ vecs.append(v / norm)
+ if not vecs:
+ continue
+ mean = np.mean(vecs, axis=0)
+ norm = float(np.linalg.norm(mean))
+ if norm > 0 and np.isfinite(norm):
+ out[speaker] = mean / norm
+ return out
+
+
+def merge_speakers_by_embedding(
+ waveform,
+ sr: int,
+ turns: list[SpeakerTurn],
+ *,
+ threshold: float = _MERGE_SIM_THRESHOLD,
+ model_name: str = DEFAULT_EMBEDDING_MODEL,
+ device: str = "auto",
+ log: LogFn = _noop,
+) -> list[SpeakerTurn]:
+ """Conservatively merge labels pyannote split for the same voice.
+
+ Builds one embedding per speaker label and unions any two labels whose cosine
+ similarity is >= *threshold*. This only fixes over-segmentation (one person
+ wrongly split into several labels); it never separates speakers. Returns the
+ (possibly relabeled) turns; on any failure it returns the input unchanged.
+ """
+ import numpy as np # noqa: PLC0415
+
+ labels = sorted({t.speaker for t in turns})
+ if len(labels) < 2:
+ return turns
+ try:
+ embs = _speaker_embeddings(
+ waveform, sr, turns, model_name=model_name, device=device, log=log
+ )
+ except Exception as exc: # noqa: BLE001
+ log(f"[diarize] speaker-embedding merge unavailable ({exc}); skipping.")
+ return turns
+
+ labels = [lab for lab in labels if lab in embs]
+ if len(labels) < 2:
+ return turns
+
+ durations: dict[str, float] = {}
+ for t in turns:
+ durations[t.speaker] = durations.get(t.speaker, 0.0) + t.duration
+
+ parent = {lab: lab for lab in labels}
+
+ def find(x: str) -> str:
+ while parent[x] != x:
+ parent[x] = parent[parent[x]]
+ x = parent[x]
+ return x
+
+ def union(a: str, b: str) -> None:
+ ra, rb = find(a), find(b)
+ if ra == rb:
+ return
+ # Keep the longer-spoken label as the group representative.
+ if durations.get(ra, 0.0) >= durations.get(rb, 0.0):
+ parent[rb] = ra
+ else:
+ parent[ra] = rb
+
+ for i in range(len(labels)):
+ for j in range(i + 1, len(labels)):
+ a, b = labels[i], labels[j]
+ sim = float(np.dot(embs[a], embs[b]))
+ if sim >= threshold:
+ log(f"[diarize] merge {a} + {b} (cosine {sim:.2f} >= {threshold:.2f}).")
+ union(a, b)
+ elif sim >= threshold - 0.1:
+ log(f"[diarize] near-miss {a} | {b} (cosine {sim:.2f} < {threshold:.2f}).")
+
+ mapping = {lab: find(lab) for lab in labels}
+ n_before = len(labels)
+ n_after = len(set(mapping.values()))
+ if n_after == n_before:
+ log(f"[diarize] no speakers merged ({n_before} kept).")
+ return turns
+
+ merged = [
+ SpeakerTurn(t.start, t.end, mapping.get(t.speaker, t.speaker)) for t in turns
+ ]
+ log(f"[diarize] merged speakers: {n_before} -> {n_after}.")
+ return merged
+
+
+# --------------------------------------------------------------------------- #
+# Cue assignment
+# --------------------------------------------------------------------------- #
+def _overlap(a_start: float, a_end: float, b_start: float, b_end: float) -> float:
+ return max(0.0, min(a_end, b_end) - max(a_start, b_start))
+
+
+def _speaker_gender(info: Any) -> str:
+ return info if isinstance(info, str) else (info or {}).get("gender", "")
+
+
+def _speaker_age_group(info: Any) -> str:
+ return "" if isinstance(info, str) else (info or {}).get("age_group", "")
+
+
+def assign_speakers_to_cues(
+ cues: list["Cue"],
+ turns: list[SpeakerTurn],
+ demographics: dict[str, Any] | None = None,
+ *,
+ nearest_gap_tolerance: float = 0.4,
+ log: LogFn = _noop,
+) -> int:
+ """Tag each cue with the speaker that overlaps it most in time.
+
+ When a cue overlaps no turn (a small diarization gap), it is assigned to the
+ nearest turn within *nearest_gap_tolerance* seconds instead of being left
+ blank. *demographics* maps a speaker label to either a gender string or a
+ dict ``{"gender","age_group",...}``. Returns the number of cues tagged.
+ """
+ demographics = demographics or {}
+ tagged = 0
+ total = len(cues)
+ log(f"[diarize] gán speaker vào {total} cue(s)...")
+ report_every = max(200, total // 20) if total else 200
+ for idx, cue in enumerate(cues, start=1):
+ best_speaker = ""
+ best_overlap = 0.0
+ nearest_speaker = ""
+ nearest_gap = float("inf")
+ for t in turns:
+ if t.end <= cue.start:
+ gap = cue.start - t.end
+ if gap < nearest_gap:
+ nearest_gap = gap
+ nearest_speaker = t.speaker
+ continue
+ if t.start >= cue.end:
+ gap = t.start - cue.end
+ if gap < nearest_gap:
+ nearest_gap = gap
+ nearest_speaker = t.speaker
+ break
+ ov = _overlap(cue.start, cue.end, t.start, t.end)
+ if ov > best_overlap:
+ best_overlap = ov
+ best_speaker = t.speaker
+ if not best_speaker and nearest_speaker and nearest_gap <= nearest_gap_tolerance:
+ best_speaker = nearest_speaker
+ if best_speaker:
+ info = demographics.get(best_speaker, {})
+ cue.speaker = best_speaker
+ cue.speaker_gender = _speaker_gender(info)
+ cue.speaker_age_group = _speaker_age_group(info)
+ tagged += 1
+ if idx % report_every == 0 or idx == total:
+ log(f"[diarize] gán cue: {idx}/{total} ({tagged} tagged)...")
+ return tagged
+
+
+# --------------------------------------------------------------------------- #
+# High level entry point
+# --------------------------------------------------------------------------- #
+def diarize_and_tag_cues(
+ video: Path,
+ cues: list["Cue"],
+ *,
+ hf_token: str | None = None,
+ num_speakers: int | None = None,
+ min_speakers: int | None = None,
+ max_speakers: int | None = None,
+ device: str = "auto",
+ detect_gender: bool = True,
+ gender_method: str = "auto",
+ age_gender_model: str = DEFAULT_AGE_GENDER_MODEL,
+ merge_speakers: bool = True,
+ merge_threshold: float = _MERGE_SIM_THRESHOLD,
+ embedding_model: str = DEFAULT_EMBEDDING_MODEL,
+ log: LogFn = _noop,
+) -> dict[str, Any]:
+ """Run the full diarization flow and tag *cues* in place.
+
+ Returns the ``{speaker: {"gender","age","age_group"}}`` registry (possibly
+ empty). Raises RuntimeError with an actionable message when the HF token is
+ missing.
+ """
+ token = resolve_hf_token(hf_token)
+ if not token:
+ raise RuntimeError(
+ "Hugging Face token required for diarization. Provide --hf-token, set "
+ "the HF_TOKEN environment variable, and accept the model conditions at "
+ f"https://hf.co/{DEFAULT_PIPELINE}"
+ )
+
+ # When the speaker count is fixed by the user, pyannote already targets it,
+ # so the embedding merge (which only fixes over-segmentation) is skipped.
+ do_merge = bool(merge_speakers) and not (num_speakers and num_speakers > 0)
+ total_steps = 4 + (1 if detect_gender else 0) + (1 if do_merge else 0)
+ step = 0
+
+ def _next(msg: str) -> None:
+ nonlocal step
+ step += 1
+ _log_step(log, step, total_steps, msg)
+
+ t_all = time.time()
+ dev = _pick_device(device)
+ n_cues = len(cues)
+ _next(
+ f"bắt đầu Pass 0 — {n_cues} cue(s), device={dev}, "
+ f"gender={gender_method if detect_gender else 'off'}"
+ )
+
+ tmp_dir = Path(tempfile.mkdtemp(prefix="gemma_diarize_"))
+ wav = tmp_dir / "audio.wav"
+ max_seconds = None
+ if cues:
+ max_seconds = max((c.end for c in cues), default=0.0) + 5.0
+ _next(f"tách audio ffmpeg (16 kHz mono, tối đa {_fmt_dur(max_seconds or 0)})...")
+ t0 = time.time()
+ extract_audio(video, wav, max_seconds=max_seconds)
+ wav_mb = wav.stat().st_size / (1024 * 1024) if wav.is_file() else 0.0
+ log(f"[diarize] audio wav ready: {wav_mb:.0f} MB in {_fmt_dur(time.time() - t0)}.")
+
+ _next("nạp waveform + chạy pyannote (who spoke when)...")
+ t0 = time.time()
+ waveform, sr = load_wav(wav)
+ aud_dur = _audio_duration_sec(waveform, sr)
+ log(f"[diarize] waveform: {_fmt_dur(aud_dur)} @ {sr} Hz.")
+
+ turns = run_diarization(
+ waveform,
+ sr,
+ token,
+ num_speakers=num_speakers,
+ min_speakers=min_speakers,
+ max_speakers=max_speakers,
+ device=device,
+ log=log,
+ )
+ log(f"[diarize] bước pyannote xong trong {_fmt_dur(time.time() - t0)}.")
+ if not turns:
+ log("[diarize] no speaker turns detected; skipping speaker tags.")
+ return {}
+
+ if do_merge:
+ n_lbl = len({t.speaker for t in turns})
+ _next(
+ f"gộp speaker bị tách nhầm bằng embedding giọng "
+ f"(ngưỡng {merge_threshold:.2f}, {n_lbl} nhãn)..."
+ )
+ t0 = time.time()
+ turns = merge_speakers_by_embedding(
+ waveform,
+ sr,
+ turns,
+ threshold=merge_threshold,
+ model_name=embedding_model,
+ device=device,
+ log=log,
+ )
+ log(f"[diarize] gộp speaker xong trong {_fmt_dur(time.time() - t0)}.")
+
+ demographics: dict[str, Any] = {}
+ if detect_gender:
+ _next(
+ f"ước lượng giới/tuổi ({gender_method}) cho "
+ f"{len({t.speaker for t in turns})} speaker(s)..."
+ )
+ t0 = time.time()
+ try:
+ demographics = estimate_speaker_demographics(
+ waveform,
+ sr,
+ turns,
+ method=gender_method,
+ model_name=age_gender_model,
+ device=device,
+ log=log,
+ )
+ log(f"[diarize] demographics xong trong {_fmt_dur(time.time() - t0)}.")
+ except Exception as exc: # noqa: BLE001
+ log(
+ f"[diarize] demographics failed ({exc}) — gán speaker không có giới/tuổi."
+ )
+
+ _next("gán speaker + giới/tuổi vào từng cue...")
+ t0 = time.time()
+ tagged = assign_speakers_to_cues(cues, turns, demographics, log=log)
+ n_spk = len({c.speaker for c in cues if c.speaker})
+ log(
+ f"[diarize] Pass 0 xong — {tagged}/{n_cues} cue tagged, "
+ f"{n_spk} speaker(s), tổng {_fmt_dur(time.time() - t_all)} "
+ f"(gán cue: {_fmt_dur(time.time() - t0)})."
+ )
+ return demographics
diff --git a/hf-upload/notebook.ipynb b/hf-upload/notebook.ipynb
new file mode 100644
index 0000000000000000000000000000000000000000..a79ccdcb33b0a6127039165aa461d24f799e001d
--- /dev/null
+++ b/hf-upload/notebook.ipynb
@@ -0,0 +1,473 @@
+{
+ "cells": [
+ {
+ "cell_type": "markdown",
+ "metadata": {},
+ "source": [
+ "# Gemma SRT Translate — Google Colab\n",
+ "\n",
+ "Dịch phụ đề SRT bằng **Gemma 4 12B vision** (llama-server + MTP).\n",
+ "\n",
+ "**Yêu cầu:** GPU **T4/L4** — Runtime → Change runtime type → GPU.\n",
+ "\n",
+ "**Lần đầu:** tải model ~12GB (+ build llama-server một lần, lưu `llama-server-bin.tgz` lên Drive). \n",
+ "**Các lần sau:** restore **`Gemma/Cache/llama-server-bin.tgz`** (~vài giây) — **không còn cell build dài**.\n",
+ "\n",
+ "---\n",
+ "\n",
+ "## Pipeline (2 bước)\n",
+ "\n",
+ "| Bước | Tên | Mặc định | Việc làm |\n",
+ "|------|-----|----------|----------|\n",
+ "| **Pass 1** | OCR / sửa SRT gốc | **Tắt** (`SKIP_CORRECTION = True`) | Cắt khung hình video → đọc sub burn-in → sửa lỗi chính tả trong `.srt` gốc |\n",
+ "| **Pass 2** | Dịch | **Luôn chạy** | Dịch sang tiếng Việt (có ngữ cảnh scene + xưng hô) |\n",
+ "\n",
+ "**`SKIP_CORRECTION` — đọc nhanh:**\n",
+ "- `True` (**mặc định**) → **bỏ Pass 1**, chỉ dịch → **nhanh hơn**\n",
+ "- `False` → **chạy Pass 1** OCR/sửa SRT trước khi dịch → **chậm hơn**, dùng khi sub gốc hay sai chữ\n",
+ "\n",
+ "**Dịch tự nhiên (lồng tiếng):** budget rộng (22 ký tự/s, tối đa ~104 ký tự/cue), **tắt Pass 2b rút gọn** mặc định.\n",
+ "\n",
+ "---\n",
+ "\n",
+ "## Cách chạy\n",
+ "\n",
+ "1. Đặt phim + SRT vào **Drive → Gemma → Phim** (cùng tên, ví dụ `Phim A.mp4` + `Phim A.srt`)\n",
+ "2. Sửa **`PHIM_STEMS`** ở cell cấu hình — tên file **không đuôi**, khớp y hệt trên Drive\n",
+ "3. Mở từ [HF /colab](https://huggingface.co/STBack23/gemma-srt-translate/colab) → **Runtime → Run all**\n",
+ "4. File `*.vi.srt` ra **Gemma/Phim** (model cache: **Gemma/Cache**)\n",
+ "\n",
+ "Sau khi xong: `DOWNLOAD_RESULTS` tải file về trình duyệt; `AUTO_UNMOUNT_DRIVE` / `AUTO_DISCONNECT_RUNTIME` tự gỡ Drive và ngắt runtime."
+ ]
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# ═══ CẤU HÌNH — sửa trước khi chạy ═══\n",
+ "from pathlib import Path\n",
+ "\n",
+ "HF_REPO = \"STBack23/gemma-srt-translate\"\n",
+ "MODELS_REPO = \"unsloth/gemma-4-12B-it-qat-GGUF\" # repo model GGUF (~12GB) — không cần sửa\n",
+ "\n",
+ "# ── Google Drive ─────────────────────────────────────────────\n",
+ "USE_DRIVE_CACHE = True # True = cache model + llama-server trên Drive (khuyến nghị)\n",
+ "\n",
+ "# Thư mục trên Drive: Gemma/Cache (model) + Gemma/Phim (video + SRT vào/ra)\n",
+ "DRIVE_ROOT = \"/content/drive/MyDrive/Gemma\"\n",
+ "DRIVE_CACHE = f\"{DRIVE_ROOT}/Cache\"\n",
+ "DRIVE_PHIM_DIR = f\"{DRIVE_ROOT}/Phim\"\n",
+ "\n",
+ "VIDEO_EXTS = (\".mp4\", \".mkv\", \".avi\", \".mov\")\n",
+ "\n",
+ "def job(stem, ext_video=\".mp4\"):\n",
+ " \"\"\"stem = tên file trên Drive (KHÔNG đuôi), khớp y hệt — kể cả dấu, khoảng trắng.\"\"\"\n",
+ " base = f\"{DRIVE_PHIM_DIR}/{stem}\"\n",
+ " return {\"stem\": stem, \"video\": f\"{base}{ext_video}\", \"srt\": f\"{base}.srt\"}\n",
+ "\n",
+ "# ── Danh sách phim cần dịch ──────────────────────────────────\n",
+ "# Tên trong PHIM_STEMS phải trùng file trong Gemma/Phim (không gồm .mp4 / .srt)\n",
+ "PHIM_STEMS = [\n",
+ " \"Thư Gửi Mùa Hè\", # → Gemma/Phim/Thư Gửi Mùa Hè.mp4 + Thư Gửi Mùa Hè.srt\n",
+ " # \"ten-phim-khac\",\n",
+ "]\n",
+ "JOBS = [job(stem) for stem in PHIM_STEMS]\n",
+ "# Cách khác — đường dẫn đầy đủ:\n",
+ "# JOBS = [{\"video\": f\"{DRIVE_PHIM_DIR}/a.mp4\", \"srt\": f\"{DRIVE_PHIM_DIR}/a.srt\"}]\n",
+ "# Upload 1 phim qua nút Colab: PHIM_STEMS = [], JOBS = [], UPLOAD_WIDGET = True\n",
+ "UPLOAD_WIDGET = False\n",
+ "\n",
+ "# ── Tùy chọn dịch ───────────────────────────────────────────\n",
+ "SOURCE_LANG = \"auto\" # ngôn ngữ SRT gốc (\"auto\" = model tự nhận)\n",
+ "TARGET_LANG = \"Vietnamese\" # ngôn ngữ đích\n",
+ "\n",
+ "# Pass 1 — OCR / sửa SRT gốc từ khung hình video (xem bảng ở đầu notebook):\n",
+ "# True → BỎ QUA Pass 1 — chỉ dịch Pass 2 (MẶC ĐỊNH, nhanh)\n",
+ "# False → CHẠY Pass 1 OCR/sửa sub trước — chậm hơn, dùng khi sub hay sai chữ\n",
+ "SKIP_CORRECTION = True\n",
+ "\n",
+ "LIMIT_CUES = 0 # 0 = dịch hết file; đặt 20 để thử nhanh vài cue đầu\n",
+ "\n",
+ "# ── Sau khi dịch xong ─────────────────────────────────────────\n",
+ "DOWNLOAD_RESULTS = True # True = tải *.vi.srt về trình duyệt\n",
+ "AUTO_UNMOUNT_DRIVE = True # True = gỡ mount Google Drive\n",
+ "AUTO_DISCONNECT_RUNTIME = True # True = ngắt runtime Colab (tiết kiệm GPU quota)\n",
+ "DISCONNECT_DELAY_SEC = 15 # giây chờ trước khi ngắt (để download kịp)\n",
+ "\n",
+ "# ── Đường dẫn nội bộ Colab (thường không cần sửa) ────────────\n",
+ "ROOT = \"/content/gemma-srt\"\n",
+ "MODELS_DIR = f\"{ROOT}/models\"\n",
+ "LLAMA_DIR = \"/content/llama.cpp\"\n",
+ "LLAMA_SERVER = f\"{LLAMA_DIR}/build/bin/llama-server\""
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# Kiểm tra GPU\n",
+ "!nvidia-smi --query-gpu=name,memory.total --format=csv,noheader\n",
+ "\n",
+ "import torch\n",
+ "if not torch.cuda.is_available():\n",
+ " raise RuntimeError(\"Chưa có GPU. Runtime → Change runtime type → GPU (L4).\")"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# Mount Google Drive — Gemma/Cache (model) + Gemma/Phim (phim)\n",
+ "if USE_DRIVE_CACHE:\n",
+ " from google.colab import drive\n",
+ " from pathlib import Path\n",
+ " import os\n",
+ " drive.mount(\"/content/drive\")\n",
+ " os.makedirs(DRIVE_CACHE, exist_ok=True)\n",
+ " os.makedirs(DRIVE_PHIM_DIR, exist_ok=True)\n",
+ " print(f\"Cache model: {DRIVE_CACHE}\")\n",
+ " print(f\"Thư mục phim: {DRIVE_PHIM_DIR}\")\n",
+ " cache_tgz = Path(f\"{DRIVE_CACHE}/llama-server-bin.tgz\")\n",
+ " if cache_tgz.is_file():\n",
+ " sz = cache_tgz.stat().st_size\n",
+ " mb = sz / 1_048_576\n",
+ " print(f\" llama-server-bin.tgz: {mb:.1f} MB\" + (\" OK\" if sz > 5_000_000 else \" HỎNG\"))\n",
+ " else:\n",
+ " for old in (\"llama-server.tgz\", \"llama-server\"):\n",
+ " p = Path(f\"{DRIVE_CACHE}/{old}\")\n",
+ " if p.is_file():\n",
+ " print(f\" cache cũ {old}: {p.stat().st_size/1024:.0f} KB — sẽ thay bằng llama-server-bin.tgz\")\n",
+ " if not any(Path(f\"{DRIVE_CACHE}/{n}\").is_file() for n in (\"llama-server-bin.tgz\", \"llama-server.tgz\", \"llama-server\")):\n",
+ " print(\" llama-server cache: chưa có (lần đầu build ~20–50 phút)\")\n",
+ " phim = Path(DRIVE_PHIM_DIR)\n",
+ " files = sorted(p.name for p in phim.iterdir() if p.is_file())\n",
+ " if files:\n",
+ " print(f\" ({len(files)} file trong Phim)\")\n",
+ " for name in files[:12]:\n",
+ " print(f\" - {name}\")\n",
+ " if len(files) > 12:\n",
+ " print(f\" ... và {len(files) - 12} file khác\")\n",
+ " else:\n",
+ " print(\" (chưa có file — upload video + .srt vào Gemma/Phim)\")"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# Cài dependency + tải code từ Hugging Face\n",
+ "!pip install -q huggingface_hub pyyaml\n",
+ "\n",
+ "from huggingface_hub import snapshot_download\n",
+ "from pathlib import Path\n",
+ "import os\n",
+ "\n",
+ "if HF_REPO.startswith(\"YOUR_\"):\n",
+ " raise ValueError(\"Sửa HF_REPO thành repo Hugging Face của bạn, ví dụ: 'username/gemma-srt-translate'\")\n",
+ "\n",
+ "print(f\"[code] Downloading {HF_REPO} ...\")\n",
+ "snapshot_download(\n",
+ " repo_id=HF_REPO,\n",
+ " repo_type=\"model\",\n",
+ " local_dir=ROOT,\n",
+ " local_dir_use_symlinks=False,\n",
+ ")\n",
+ "print(f\"[code] OK: {ROOT}\")\n",
+ "assert Path(f\"{ROOT}/translate_srt.py\").is_file(), \"translate_srt.py not found in repo\""
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# Tải model GGUF (~12GB) — cache trên Drive nếu bật\n",
+ "from huggingface_hub import hf_hub_download\n",
+ "from pathlib import Path\n",
+ "import os\n",
+ "import shutil\n",
+ "\n",
+ "MODEL_FILES = [\n",
+ " \"gemma-4-12B-it-qat-UD-Q4_K_XL.gguf\",\n",
+ " \"mmproj-F16.gguf\",\n",
+ " \"mtp-gemma-4-12B-it.gguf\",\n",
+ "]\n",
+ "\n",
+ "models_dest = MODELS_DIR\n",
+ "if USE_DRIVE_CACHE:\n",
+ " models_dest = f\"{DRIVE_CACHE}/models\"\n",
+ "Path(models_dest).mkdir(parents=True, exist_ok=True)\n",
+ "\n",
+ "paths = {}\n",
+ "for name in MODEL_FILES:\n",
+ " dest_file = Path(models_dest) / name\n",
+ " if dest_file.is_file():\n",
+ " print(f\"[model] cached: {name}\")\n",
+ " paths[name] = dest_file\n",
+ " continue\n",
+ " print(f\"[model] downloading {MODELS_REPO}/{name} ...\")\n",
+ " p = hf_hub_download(\n",
+ " repo_id=MODELS_REPO,\n",
+ " filename=name,\n",
+ " local_dir=models_dest,\n",
+ " local_dir_use_symlinks=False,\n",
+ " )\n",
+ " paths[name] = Path(p)\n",
+ " print(f\" OK: {p}\")\n",
+ "\n",
+ "# Symlink/copy vào ROOT/models cho translate_srt.py\n",
+ "Path(MODELS_DIR).mkdir(parents=True, exist_ok=True)\n",
+ "for name, src in paths.items():\n",
+ " dst = Path(MODELS_DIR) / name\n",
+ " if not dst.exists():\n",
+ " try:\n",
+ " os.symlink(src, dst)\n",
+ " except OSError:\n",
+ " shutil.copy2(src, dst)\n",
+ "\n",
+ "MODEL_PATH = Path(MODELS_DIR) / MODEL_FILES[0]\n",
+ "MMPROJ_PATH = Path(MODELS_DIR) / MODEL_FILES[1]\n",
+ "DRAFT_PATH = Path(MODELS_DIR) / MODEL_FILES[2]\n",
+ "print(f\"\\nModels ready in {MODELS_DIR}\")"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# llama-server — restore từ Drive (Gemma/Cache/llama-server-bin.tgz)\n",
+ "import sys\n",
+ "\n",
+ "sys.path.insert(0, ROOT)\n",
+ "from scripts.ensure_llama_colab import ensure_llama_server\n",
+ "\n",
+ "LLAMA_SERVER = str(ensure_llama_server(\n",
+ " llama_dir=LLAMA_DIR,\n",
+ " drive_cache=DRIVE_CACHE if USE_DRIVE_CACHE else None,\n",
+ " allow_build=False,\n",
+ "))\n",
+ "print(f\"LLAMA_SERVER = {LLAMA_SERVER}\")\n"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# Kiểm tra ffmpeg\n",
+ "!ffmpeg -version | head -1"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# Kiểm tra danh sách JOBS (Drive) hoặc upload 1 phim qua widget\n",
+ "from pathlib import Path\n",
+ "from google.colab import files\n",
+ "import os\n",
+ "\n",
+ "UPLOAD_DIR = \"/content/uploads\"\n",
+ "os.makedirs(UPLOAD_DIR, exist_ok=True)\n",
+ "resolved_jobs = []\n",
+ "\n",
+ "def _list_phim_hint() -> str:\n",
+ " phim = Path(DRIVE_PHIM_DIR)\n",
+ " if not phim.is_dir():\n",
+ " return f\"Chưa có thư mục: {DRIVE_PHIM_DIR}\"\n",
+ " lines = [\"File hiện có trong Gemma/Phim:\"]\n",
+ " for p in sorted(phim.iterdir()):\n",
+ " if p.is_file():\n",
+ " lines.append(f\" - {p.name}\")\n",
+ " if len(lines) == 1:\n",
+ " lines.append(\" (trống — upload video + .srt vào đây)\")\n",
+ " return \"\\n\".join(lines)\n",
+ "\n",
+ "def _find_video(stem: str, preferred: str | None = None) -> Path | None:\n",
+ " if preferred:\n",
+ " p = Path(preferred)\n",
+ " if p.is_file():\n",
+ " return p\n",
+ " for ext in VIDEO_EXTS:\n",
+ " p = Path(DRIVE_PHIM_DIR) / f\"{stem}{ext}\"\n",
+ " if p.is_file():\n",
+ " return p\n",
+ " return None\n",
+ "\n",
+ "if JOBS:\n",
+ " for i, entry in enumerate(JOBS, 1):\n",
+ " stem = entry.get(\"stem\") or Path(entry[\"video\"]).stem\n",
+ " v = _find_video(stem, entry.get(\"video\"))\n",
+ " s = Path(entry[\"srt\"])\n",
+ " if v is None:\n",
+ " raise FileNotFoundError(\n",
+ " f\"Job {i}: không thấy video cho '{stem}'\\n\"\n",
+ " f\"Đã thử: {', '.join(stem + ext for ext in VIDEO_EXTS)}\\n\\n\"\n",
+ " f\"Tên trong PHIM_STEMS phải KHỚP Y HỆT tên file (không đuôi).\\n\\n\"\n",
+ " f\"{_list_phim_hint()}\"\n",
+ " )\n",
+ " if not s.is_file():\n",
+ " s = Path(DRIVE_PHIM_DIR) / f\"{stem}.srt\"\n",
+ " if not s.is_file():\n",
+ " raise FileNotFoundError(\n",
+ " f\"Job {i}: không thấy SRT: {s}\\n\\n\"\n",
+ " f\"Cần file: {stem}.srt cùng thư mục Phim.\\n\\n\"\n",
+ " f\"{_list_phim_hint()}\"\n",
+ " )\n",
+ " out = Path(entry[\"output\"]) if entry.get(\"output\") else v.with_name(v.stem + \".vi.srt\")\n",
+ " resolved_jobs.append({\"video\": v, \"srt\": s, \"output\": out})\n",
+ " print(f\"Job {i}/{len(JOBS)}: {v.name} + {s.name} -> {out.name}\")\n",
+ "elif UPLOAD_WIDGET:\n",
+ " print(\"Upload 1 video (.mp4/.mkv) và 1 SRT (.srt):\")\n",
+ " uploaded = files.upload()\n",
+ " video = srt = None\n",
+ " for name in uploaded:\n",
+ " p = Path(UPLOAD_DIR) / name\n",
+ " p.write_bytes(uploaded[name])\n",
+ " low = name.lower()\n",
+ " if low.endswith((\".mp4\", \".mkv\", \".avi\", \".mov\")):\n",
+ " video = p\n",
+ " elif low.endswith(\".srt\"):\n",
+ " srt = p\n",
+ " if not video or not srt:\n",
+ " raise ValueError(\"Cần 1 file video và 1 file .srt\")\n",
+ " out = video.with_name(video.stem + \".vi.srt\")\n",
+ " resolved_jobs.append({\"video\": video, \"srt\": srt, \"output\": out})\n",
+ " print(f\"Job 1/1: {video.name} -> {out.name}\")\n",
+ "else:\n",
+ " raise ValueError(\"Thêm phim vào JOBS hoặc đặt UPLOAD_WIDGET = True\")\n",
+ "\n",
+ "print(f\"\\nTổng: {len(resolved_jobs)} phim (chạy tuần tự)\")"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# ═══ CHẠY DỊCH SRT (tuần tự từng phim) ═══\n",
+ "import sys\n",
+ "import os\n",
+ "from pathlib import Path\n",
+ "\n",
+ "sys.path.insert(0, ROOT)\n",
+ "from scripts.ensure_llama_colab import ensure_llama_server, llama_bin_valid\n",
+ "\n",
+ "if not llama_bin_valid(Path(LLAMA_SERVER)):\n",
+ " LLAMA_SERVER = str(ensure_llama_server(\n",
+ " llama_dir=LLAMA_DIR,\n",
+ " drive_cache=DRIVE_CACHE if USE_DRIVE_CACHE else None,\n",
+ " allow_build=False,\n",
+ " ))\n",
+ "\n",
+ "from translate_srt import build_parser, run_pipeline\n",
+ "\n",
+ "completed = []\n",
+ "\n",
+ "for i, job in enumerate(resolved_jobs, 1):\n",
+ " v, s, out = job[\"video\"], job[\"srt\"], job[\"output\"]\n",
+ " print(\"\\n\" + \"=\" * 50)\n",
+ " print(f\"=== Phim {i}/{len(resolved_jobs)}: {v.name} ===\")\n",
+ " print(\"=\" * 50)\n",
+ "\n",
+ " argv = [\n",
+ " \"--video\", str(v),\n",
+ " \"--input-srt\", str(s),\n",
+ " \"--output-srt\", str(out),\n",
+ " \"--source-lang\", SOURCE_LANG,\n",
+ " \"--target-lang\", TARGET_LANG,\n",
+ " \"--llama-server\", LLAMA_SERVER,\n",
+ " \"--model\", str(MODEL_PATH),\n",
+ " \"--mmproj\", str(MMPROJ_PATH),\n",
+ " \"--model-draft\", str(DRAFT_PATH),\n",
+ " \"--limit\", str(LIMIT_CUES),\n",
+ " \"--ngl\", \"999\",\n",
+ " \"--ctx\", \"8192\",\n",
+ " ]\n",
+ " if SKIP_CORRECTION:\n",
+ " argv.append(\"--skip-correction\")\n",
+ "\n",
+ " args = build_parser().parse_args(argv)\n",
+ " args.log_fn = lambda msg, _i=i: print(f\"[{_i}] {msg}\", flush=True)\n",
+ " run_pipeline(args)\n",
+ " completed.append(out)\n",
+ " print(f\"\\nXong phim {i}: {out}\")\n",
+ "\n",
+ "print(\"\\n\" + \"=\" * 50)\n",
+ "print(f\"HOÀN TẤT {len(completed)} phim:\")\n",
+ "for p in completed:\n",
+ " print(f\" - {p}\")"
+ ],
+ "execution_count": null,
+ "outputs": []
+ },
+ {
+ "cell_type": "code",
+ "metadata": {},
+ "source": [
+ "# Tải kết quả về máy (tuỳ chọn) + ngắt kết nối\n",
+ "from pathlib import Path\n",
+ "import time\n",
+ "\n",
+ "if DOWNLOAD_RESULTS:\n",
+ " from google.colab import files\n",
+ " for out in completed:\n",
+ " out = Path(out)\n",
+ " for p in [\n",
+ " out,\n",
+ " out.with_suffix(\".corrected\" + out.suffix),\n",
+ " out.with_suffix(out.suffix + \".report.json\"),\n",
+ " ]:\n",
+ " if p.is_file():\n",
+ " print(f\"Download: {p.name}\")\n",
+ " files.download(str(p))\n",
+ "else:\n",
+ " print(\"Bỏ qua download — file SRT đã nằm trên Drive (xem đường dẫn output trong JOBS).\")\n",
+ "\n",
+ "if AUTO_UNMOUNT_DRIVE and USE_DRIVE_CACHE:\n",
+ " from google.colab import drive\n",
+ " drive.flush_and_unmount()\n",
+ " print(\"Đã gỡ mount Google Drive.\")\n",
+ "\n",
+ "if AUTO_DISCONNECT_RUNTIME:\n",
+ " delay = max(0, int(DISCONNECT_DELAY_SEC))\n",
+ " if delay:\n",
+ " print(f\"Ngắt Colab runtime sau {delay}s...\")\n",
+ " time.sleep(delay)\n",
+ " from google.colab import runtime\n",
+ " runtime.unassign()"
+ ],
+ "execution_count": null,
+ "outputs": []
+ }
+ ],
+ "metadata": {
+ "accelerator": "GPU",
+ "colab": {
+ "gpuType": "L4",
+ "provenance": []
+ },
+ "kernelspec": {
+ "display_name": "Python 3",
+ "name": "python3"
+ },
+ "language_info": {
+ "name": "python"
+ }
+ },
+ "nbformat": 4,
+ "nbformat_minor": 0
+}
\ No newline at end of file
diff --git a/hf-upload/requirements-diarize.txt b/hf-upload/requirements-diarize.txt
new file mode 100644
index 0000000000000000000000000000000000000000..346f44ae5a350c0ca6c26b902b0d569a34f7d5cf
--- /dev/null
+++ b/hf-upload/requirements-diarize.txt
@@ -0,0 +1,24 @@
+# Optional dependencies for speaker diarization (Pass 0).
+# Only needed when you enable "Phân tích giọng nói nhân vật" / --diarize.
+#
+# Setup (Windows, NVIDIA GPU recommended):
+# 1. Install PyTorch with CUDA matching your driver, e.g.:
+# pip install torch torchaudio --index-url https://download.pytorch.org/whl/cu124
+# (CPU-only also works but is much slower:)
+# pip install torch torchaudio
+# 2. pip install -r requirements-diarize.txt
+# 3. Create a Hugging Face token at https://hf.co/settings/tokens
+# 4. Accept the model conditions at:
+# https://hf.co/pyannote/speaker-diarization-community-1
+# 5. Set HF_TOKEN env var or paste the token into the app's "HF token" field.
+#
+# torchcodec is used by pyannote.audio 4.x for audio decoding and needs ffmpeg
+# available on PATH (this project already ships/uses ffmpeg).
+
+# Audio is decoded with soundfile (not torchcodec) for a robust Windows setup.
+pyannote.audio>=4.0
+torch
+torchaudio
+soundfile>=0.11
+# transformers is needed for the audeering age/gender model (--gender-method model).
+transformers>=4.40
diff --git a/hf-upload/scripts/build_llama_server.sh b/hf-upload/scripts/build_llama_server.sh
new file mode 100644
index 0000000000000000000000000000000000000000..a755553e4b49326a4bce6cef569732a8c3a67528
--- /dev/null
+++ b/hf-upload/scripts/build_llama_server.sh
@@ -0,0 +1,38 @@
+#!/usr/bin/env bash
+# Build llama-server with CUDA for Colab (T4=75, L4=89).
+set -euo pipefail
+
+LLAMA_DIR="${1:-/content/llama.cpp}"
+CUDA_ARCH="${2:-75}"
+
+export CUDA_HOME="${CUDA_HOME:-/usr/local/cuda}"
+export PATH="${CUDA_HOME}/bin:${PATH}"
+export LD_LIBRARY_PATH="${CUDA_HOME}/lib64:${LD_LIBRARY_PATH:-}"
+
+echo "[build] CUDA_HOME=${CUDA_HOME} arch=${CUDA_ARCH}"
+
+if [[ ! -d "${LLAMA_DIR}/.git" ]]; then
+ git clone --depth 1 https://github.com/ggml-org/llama.cpp.git "${LLAMA_DIR}"
+else
+ git -C "${LLAMA_DIR}" fetch origin master
+ git -C "${LLAMA_DIR}" reset --hard origin/master
+fi
+
+rm -rf "${LLAMA_DIR}/build"
+
+cmake -S "${LLAMA_DIR}" -B "${LLAMA_DIR}/build" \
+ -DGGML_CUDA=ON \
+ -DCMAKE_CUDA_ARCHITECTURES="${CUDA_ARCH}" \
+ -DCMAKE_BUILD_TYPE=Release \
+ -DCMAKE_CUDA_COMPILER="${CUDA_HOME}/bin/nvcc"
+
+cmake --build "${LLAMA_DIR}/build" --config Release -j "$(nproc)" --target llama-server
+
+BIN="${LLAMA_DIR}/build/bin/llama-server"
+if [[ ! -x "${BIN}" ]]; then
+ echo "[build] ERROR: ${BIN} not found"
+ exit 1
+fi
+
+echo "[build] OK: ${BIN}"
+"${BIN}" --version 2>/dev/null || true
diff --git a/hf-upload/scripts/download_models.py b/hf-upload/scripts/download_models.py
new file mode 100644
index 0000000000000000000000000000000000000000..7d828748f889262b1306c877a6037b633a912fbe
--- /dev/null
+++ b/hf-upload/scripts/download_models.py
@@ -0,0 +1,72 @@
+#!/usr/bin/env python3
+"""Tải model Gemma 4 QAT + mmproj + MTP từ Hugging Face."""
+
+from __future__ import annotations
+
+import argparse
+from pathlib import Path
+
+DEFAULT_SOURCE = "unsloth/gemma-4-12B-it-qat-GGUF"
+DEFAULT_FILES = [
+ "gemma-4-12B-it-qat-UD-Q4_K_XL.gguf",
+ "mmproj-F16.gguf",
+ "mtp-gemma-4-12B-it.gguf",
+]
+
+
+def download(
+ dest: Path,
+ source_repo: str = DEFAULT_SOURCE,
+ files: list[str] | None = None,
+ token: str | None = None,
+) -> dict[str, Path]:
+ from huggingface_hub import hf_hub_download
+
+ dest.mkdir(parents=True, exist_ok=True)
+ files = files or DEFAULT_FILES
+ result: dict[str, Path] = {}
+
+ for name in files:
+ print(f"[download] {source_repo}/{name} -> {dest}/")
+ path = hf_hub_download(
+ repo_id=source_repo,
+ filename=name,
+ local_dir=str(dest),
+ local_dir_use_symlinks=False,
+ token=token,
+ )
+ result[name] = Path(path)
+ print(f" OK: {path}")
+
+ return result
+
+
+def main() -> int:
+ p = argparse.ArgumentParser(description="Download Gemma 4 GGUF files for SRT translation")
+ p.add_argument(
+ "--dest",
+ type=Path,
+ default=Path("models"),
+ help="Output directory (default: ./models)",
+ )
+ p.add_argument(
+ "--source-repo",
+ default=DEFAULT_SOURCE,
+ help=f"Hugging Face model repo (default: {DEFAULT_SOURCE})",
+ )
+ p.add_argument(
+ "--file",
+ action="append",
+ dest="files",
+ help="Download specific file(s); repeat flag. Default: all 3 core files.",
+ )
+ p.add_argument("--token", default=None, help="HF token (optional, for gated models)")
+ args = p.parse_args()
+
+ download(args.dest, args.source_repo, args.files, args.token)
+ print(f"\nDone. Models in: {args.dest.resolve()}")
+ return 0
+
+
+if __name__ == "__main__":
+ raise SystemExit(main())
diff --git a/hf-upload/scripts/ensure_llama_colab.py b/hf-upload/scripts/ensure_llama_colab.py
new file mode 100644
index 0000000000000000000000000000000000000000..ca9deb7b8b76ba96bfa7bb308f0168d27712e22e
--- /dev/null
+++ b/hf-upload/scripts/ensure_llama_colab.py
@@ -0,0 +1,441 @@
+"""Restore or build llama-server for Google Colab (CUDA + Drive cache)."""
+
+from __future__ import annotations
+
+import os
+import re
+import shutil
+import subprocess
+import tarfile
+import time
+from pathlib import Path
+
+MIN_LLAMA_TGZ_BYTES = 5_000_000
+DRIVE_LLAMA_TGZ = "llama-server-bin.tgz"
+DRIVE_COPY_BLOCK = 4 * 1024 * 1024
+# Gemma 4 draft-MTP merged 2026-06-07 (llama.cpp PR #23398).
+MIN_LLAMA_BUILD_FOR_MTP = 9549
+# Tag b9553 (9e3b928) — verified Gemma 4 draft-mtp on L4; newer master may regress.
+LLAMA_CPP_PIN = "b9553"
+
+
+def _fmt_size(n: int) -> str:
+ if n >= 1_048_576:
+ return f"{n / 1_048_576:.1f} MB"
+ return f"{n / 1024:.0f} KB"
+
+
+def _run(cmd, *, check=True, label="", live=False):
+ print(f"[llama] $ {' '.join(map(str, cmd))}", flush=True)
+ if live:
+ proc = subprocess.Popen(
+ cmd, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True, bufsize=1,
+ )
+ assert proc.stdout is not None
+ for line in proc.stdout:
+ print(line, end="", flush=True)
+ proc.wait()
+ if check and proc.returncode:
+ raise RuntimeError(f"{label or cmd[0]} failed (exit {proc.returncode})")
+ return proc
+ r = subprocess.run(cmd, capture_output=True, text=True)
+ if r.stdout:
+ print(r.stdout[-4000:])
+ if r.returncode and r.stderr:
+ print("--- stderr ---")
+ print(r.stderr[-8000:])
+ if check and r.returncode:
+ raise RuntimeError(f"{label or cmd[0]} failed (exit {r.returncode})")
+ return r
+
+
+def _setup_cuda_env() -> Path | None:
+ for cuda in (Path("/usr/local/cuda"), Path("/usr/local/cuda-12.2"), Path("/usr/local/cuda-12.4")):
+ nvcc = cuda / "bin" / "nvcc"
+ if nvcc.is_file():
+ os.environ["CUDA_HOME"] = str(cuda)
+ os.environ["PATH"] = f"{cuda / 'bin'}:" + os.environ.get("PATH", "")
+ os.environ["LD_LIBRARY_PATH"] = f"{cuda / 'lib64'}:" + os.environ.get("LD_LIBRARY_PATH", "")
+ print(f"[llama] CUDA: {cuda} | nvcc OK")
+ return nvcc
+ return None
+
+
+def _cuda_arch() -> str:
+ try:
+ out = subprocess.check_output(
+ ["nvidia-smi", "--query-gpu=compute_cap", "--format=csv,noheader"],
+ text=True, timeout=10,
+ ).strip().splitlines()[0]
+ return out.replace(".", "")
+ except Exception:
+ return "75"
+
+
+def _llama_env_for(p: Path) -> dict:
+ env = os.environ.copy()
+ lib = str(p.parent)
+ env["LD_LIBRARY_PATH"] = f"{lib}:{env.get('LD_LIBRARY_PATH', '')}"
+ return env
+
+
+def llama_version_text(p: Path) -> str:
+ p = Path(p)
+ try:
+ r = subprocess.run(
+ [str(p), "--version"],
+ capture_output=True, text=True, timeout=60, env=_llama_env_for(p),
+ )
+ except Exception:
+ return ""
+ return ((r.stdout or "") + (r.stderr or "")).strip()
+
+
+def llama_build_number(p: Path) -> int | None:
+ """Legacy llama.cpp prints ``version: 9535 (hash)``; newer builds may use semver."""
+ text = llama_version_text(p)
+ m = re.search(r"version:\s*(\d+)", text)
+ if not m:
+ return None
+ n = int(m.group(1))
+ # Newer releases use small semver (e.g. version: 1) — not comparable to b9549.
+ return n if n >= 1000 else None
+
+
+def _llama_help_text(p: Path) -> str:
+ try:
+ r = subprocess.run(
+ [str(p), "--help"],
+ capture_output=True,
+ text=True,
+ timeout=120,
+ env=_llama_env_for(p),
+ )
+ except Exception:
+ return ""
+ return (r.stdout or "") + (r.stderr or "")
+
+
+def llama_supports_mtp(p: Path) -> bool:
+ if not llama_bin_valid(p):
+ return False
+ if "draft-mtp" in _llama_help_text(p):
+ return True
+ n = llama_build_number(p)
+ return n is not None and n >= MIN_LLAMA_BUILD_FOR_MTP
+
+
+def llama_bin_valid(p: Path) -> bool:
+ """Shared build: llama-server ~17–20 KB + .so cùng thư mục."""
+ p = Path(p)
+ if not p.is_file() or not os.access(p, os.X_OK):
+ return False
+ if p.stat().st_size < 8_000:
+ return False
+ return bool(llama_version_text(p))
+
+
+def setup_ld_library_path(bin_path: Path) -> None:
+ lib = str(bin_path.parent)
+ os.environ["LD_LIBRARY_PATH"] = f"{lib}:{os.environ.get('LD_LIBRARY_PATH', '')}"
+
+
+def _print_bin_info(bin_path: Path, *, require_mtp: bool) -> None:
+ ver = llama_version_text(bin_path)
+ build = llama_build_number(bin_path)
+ mtp = llama_supports_mtp(bin_path)
+ print(f"[llama] OK: {bin_path} ({_fmt_size(bin_path.stat().st_size)})")
+ if ver:
+ print(f"[llama] {ver.splitlines()[0]}")
+ if require_mtp:
+ if mtp:
+ tag = f"build {build}" if build else "draft-mtp in --help"
+ print(f"[llama] MTP Gemma 4: supported ({tag})")
+ else:
+ print("[llama] MTP Gemma 4: NOT supported — will rebuild")
+
+
+def _purge_invalid_llama_bin(build_dir: Path) -> bool:
+ bin_dir = build_dir / "bin"
+ if not bin_dir.is_dir():
+ return False
+ p = bin_dir / "llama-server"
+ if p.is_file() and not llama_bin_valid(p):
+ print(f"[llama] Xóa binary hỏng (cmake sẽ link lại): {p} ({_fmt_size(p.stat().st_size)})")
+ p.unlink(missing_ok=True)
+ return True
+ return False
+
+
+def _invalidate_llama_artifacts(llama_dir: Path, drive_cache: Path | None) -> None:
+ build_dir = llama_dir / "build"
+ if build_dir.exists():
+ print(f"[llama] Xóa build cũ: {build_dir}")
+ shutil.rmtree(build_dir, ignore_errors=True)
+ if drive_cache:
+ drive_tgz = drive_cache / DRIVE_LLAMA_TGZ
+ if drive_tgz.is_file():
+ print(f"[llama] Xóa cache Drive cũ: {drive_tgz}")
+ drive_tgz.unlink(missing_ok=True)
+ for legacy_name in ("llama-server", "llama-server.tgz"):
+ legacy = drive_cache / legacy_name
+ if legacy.is_file():
+ legacy.unlink(missing_ok=True)
+
+
+def _drive_copy_file(src: Path, dest: Path, *, label: str = "", min_bytes: int = 1) -> int:
+ src = Path(src)
+ dest = Path(dest)
+ dest.parent.mkdir(parents=True, exist_ok=True)
+ sz_src = src.stat().st_size
+ if dest.exists():
+ dest.unlink()
+ tmp = dest.with_name(dest.name + ".part")
+ if tmp.exists():
+ tmp.unlink()
+ with open(src, "rb") as fsrc, open(tmp, "wb") as fdst:
+ while True:
+ chunk = fsrc.read(DRIVE_COPY_BLOCK)
+ if not chunk:
+ break
+ fdst.write(chunk)
+ fdst.flush()
+ os.fsync(fdst.fileno())
+ sz_tmp = tmp.stat().st_size
+ if sz_tmp != sz_src or sz_tmp < min_bytes:
+ tmp.unlink(missing_ok=True)
+ raise RuntimeError(f"{label}copy thất bại: {sz_tmp} bytes, cần {sz_src}")
+ tmp.replace(dest)
+ dest.chmod(0o755)
+ for _ in range(5):
+ sz_dest = dest.stat().st_size
+ if sz_dest == sz_src:
+ return sz_dest
+ time.sleep(1)
+ raise RuntimeError(
+ f"{label}Drive vẫn sai size sau copy: {dest.stat().st_size} bytes, cần {sz_src}"
+ )
+
+
+def _find_llama_bin(llama_dir: Path, llama_server: Path | None) -> Path | None:
+ candidates = []
+ if llama_server:
+ candidates.append(Path(llama_server))
+ candidates.extend([
+ llama_dir / "build/bin/llama-server",
+ llama_dir / "build/llama-server",
+ ])
+ seen: set[Path] = set()
+ for p in candidates:
+ p = p.resolve()
+ if p in seen:
+ continue
+ seen.add(p)
+ if llama_bin_valid(p):
+ return p
+ if p.is_file():
+ print(f"[llama] Bỏ qua (không chạy --version): {p} ({_fmt_size(p.stat().st_size)})")
+ return None
+
+
+def _restore_llama_from_drive_tgz(tgz: Path, dest: Path) -> Path:
+ local_tgz = Path("/content/llama-server-restore.tgz")
+ _drive_copy_file(tgz, local_tgz, label="restore tgz ", min_bytes=MIN_LLAMA_TGZ_BYTES)
+ if dest.parent.exists():
+ for old in dest.parent.iterdir():
+ if old.is_file():
+ old.unlink()
+ dest.parent.mkdir(parents=True, exist_ok=True)
+ with tarfile.open(local_tgz, "r:gz") as tar:
+ tar.extractall(path=dest.parent)
+ dest.chmod(0o755)
+ if not llama_bin_valid(dest):
+ raise RuntimeError("Giải nén cache thất bại — thử build lại")
+ local_tgz.unlink(missing_ok=True)
+ n = sum(1 for f in dest.parent.iterdir() if f.is_file())
+ print(f"[llama] restored {n} file từ Drive → {dest.parent}")
+ return dest
+
+
+def _cache_llama_to_drive(bin_path: Path, drive_cache: Path) -> None:
+ bin_dir = bin_path.parent
+ local_tgz = Path("/content/llama-server-cache.tgz")
+ if local_tgz.exists():
+ local_tgz.unlink()
+ with tarfile.open(local_tgz, "w:gz") as tar:
+ for f in sorted(bin_dir.iterdir()):
+ if f.is_file():
+ tar.add(f, arcname=f.name)
+ drive_tgz = drive_cache / DRIVE_LLAMA_TGZ
+ sz_tgz = _drive_copy_file(local_tgz, drive_tgz, label="cache tgz ", min_bytes=MIN_LLAMA_TGZ_BYTES)
+ local_tgz.unlink(missing_ok=True)
+ for legacy_name in ("llama-server", "llama-server.tgz"):
+ legacy = drive_cache / legacy_name
+ if legacy.is_file():
+ legacy.unlink(missing_ok=True)
+ build = llama_build_number(bin_path)
+ print(
+ f"[llama] cached to Drive: {drive_tgz} "
+ f"(build {build or '?'}, llama-server {_fmt_size(bin_path.stat().st_size)}, "
+ f"tgz {_fmt_size(sz_tgz)}, {sum(1 for f in bin_dir.iterdir() if f.is_file())} files)"
+ )
+
+
+def _sync_llama_source(llama_dir: Path) -> None:
+ repo = "https://github.com/ggml-org/llama.cpp.git"
+ tag = LLAMA_CPP_PIN
+ if llama_dir.exists():
+ shutil.rmtree(llama_dir, ignore_errors=True)
+ llama_dir.parent.mkdir(parents=True, exist_ok=True)
+ _run(
+ ["git", "clone", "--depth", "1", "--branch", tag, repo, str(llama_dir)],
+ label=f"git clone tag {tag}",
+ )
+
+
+def _build_llama_server(
+ llama_dir: Path,
+ drive_cache: Path | None,
+ *,
+ fresh_source: bool = False,
+) -> Path:
+ cuda_arch = _cuda_arch()
+ nproc = min(4, os.cpu_count() or 2)
+ print(
+ f"[llama] Building llama-server (CUDA arch {cuda_arch}, -j {nproc}, "
+ f"pin {LLAMA_CPP_PIN}, MTP min build {MIN_LLAMA_BUILD_FOR_MTP})...",
+ flush=True,
+ )
+ print("[llama] Compile L4/T4: thường 20–50 phút. Log in từng dòng — RAM GPU ~0 là bình thường.", flush=True)
+
+ _run(["apt-get", "-qq", "update"], check=False)
+ _run([
+ "apt-get", "-qq", "install", "-y",
+ "build-essential", "cmake", "git", "pkg-config", "libcurl4-openssl-dev",
+ ], label="apt build deps")
+
+ nvcc = _setup_cuda_env()
+ if nvcc is None:
+ for pkg in ("cuda-nvcc-12-2", "cuda-nvcc-12-4", "nvidia-cuda-toolkit"):
+ r = _run(["apt-get", "-qq", "install", "-y", pkg], check=False, label=pkg)
+ if r.returncode == 0:
+ nvcc = _setup_cuda_env()
+ if nvcc:
+ break
+ if nvcc is None:
+ raise RuntimeError(
+ "Không tìm thấy nvcc. Runtime → Change runtime type → GPU (T4/L4), Restart, Run all."
+ )
+
+ if fresh_source:
+ _sync_llama_source(llama_dir)
+
+ build_dir = llama_dir / "build"
+ if fresh_source and build_dir.exists():
+ shutil.rmtree(build_dir, ignore_errors=True)
+
+ cache_file = build_dir / "CMakeCache.txt"
+ if not cache_file.is_file():
+ if build_dir.exists():
+ shutil.rmtree(build_dir, ignore_errors=True)
+ _run([
+ "cmake", "-S", str(llama_dir), "-B", str(build_dir),
+ "-DGGML_CUDA=ON",
+ f"-DCMAKE_CUDA_ARCHITECTURES={cuda_arch}",
+ "-DCMAKE_BUILD_TYPE=Release",
+ f"-DCMAKE_CUDA_COMPILER={nvcc}",
+ "-DLLAMA_BUILD_TESTS=OFF",
+ "-DLLAMA_BUILD_EXAMPLES=OFF",
+ ], label="cmake configure")
+ else:
+ print("[llama] Tiếp tục build cũ (CMakeCache.txt có sẵn)...", flush=True)
+
+ need_clean = _purge_invalid_llama_bin(build_dir)
+ print("[llama] Đang compile llama-server...", flush=True)
+ build_cmd = [
+ "cmake", "--build", str(build_dir), "--config", "Release",
+ "-j", str(nproc), "--target", "llama-server",
+ ]
+ if need_clean:
+ build_cmd.insert(4, "--clean-first")
+ print("[llama] --clean-first (xóa stub hỏng, buộc link lại)", flush=True)
+ _run(build_cmd, label="cmake build", live=True)
+
+ bin_path = _find_llama_bin(llama_dir, None)
+ if not bin_path:
+ raise FileNotFoundError(
+ f"Build xong nhưng llama-server không chạy được trong {build_dir}/bin."
+ )
+ if not llama_supports_mtp(bin_path):
+ raise RuntimeError(
+ "Build xong nhưng llama-server không hỗ trợ draft-mtp (MTP Gemma 4). "
+ "Thử FORCE_LLAMA_REBUILD = True."
+ )
+ print(f"[llama] Build OK: {bin_path} ({_fmt_size(bin_path.stat().st_size)} + .so cùng thư mục)")
+ print(llama_version_text(bin_path).splitlines()[0] if llama_version_text(bin_path) else str(bin_path))
+ if drive_cache:
+ _cache_llama_to_drive(bin_path, drive_cache)
+ return bin_path
+
+
+def ensure_llama_server(
+ *,
+ llama_dir: str | Path = "/content/llama.cpp",
+ llama_server: str | Path | None = None,
+ drive_cache: str | Path | None = None,
+ allow_build: bool = True,
+ require_mtp: bool = True,
+ force_rebuild: bool = False,
+) -> Path:
+ """Restore từ Drive cache hoặc build llama.cpp mới (>= b9549 cho MTP Gemma 4)."""
+ llama_dir = Path(llama_dir)
+ llama_server = Path(llama_server) if llama_server else llama_dir / "build/bin/llama-server"
+ drive_cache = Path(drive_cache) if drive_cache else None
+
+ if force_rebuild:
+ print("[llama] FORCE_LLAMA_REBUILD=True — xóa cache + build lại...")
+ _invalidate_llama_artifacts(llama_dir, drive_cache)
+ elif require_mtp:
+ bin_path = _find_llama_bin(llama_dir, llama_server)
+ if bin_path and not llama_supports_mtp(bin_path):
+ print("[llama] Binary cũ (không có draft-mtp) — xóa và build lại...")
+ _invalidate_llama_artifacts(llama_dir, drive_cache)
+
+ bin_path = _find_llama_bin(llama_dir, llama_server)
+ if bin_path and (not require_mtp or llama_supports_mtp(bin_path)):
+ _print_bin_info(bin_path, require_mtp=require_mtp)
+ setup_ld_library_path(bin_path)
+ return bin_path
+
+ if drive_cache and not force_rebuild:
+ drive_tgz = drive_cache / DRIVE_LLAMA_TGZ
+ if drive_tgz.is_file() and drive_tgz.stat().st_size > MIN_LLAMA_TGZ_BYTES:
+ dest = llama_dir / "build/bin/llama-server"
+ try:
+ _restore_llama_from_drive_tgz(drive_tgz, dest)
+ if not require_mtp or llama_supports_mtp(dest):
+ _print_bin_info(dest, require_mtp=require_mtp)
+ setup_ld_library_path(dest)
+ return dest
+ print("[llama] Cache Drive không hỗ trợ draft-mtp — build lại cho MTP...")
+ _invalidate_llama_artifacts(llama_dir, drive_cache)
+ except Exception as e:
+ print(f"[llama] Cache .tgz hỏng ({e}) — build lại...")
+ drive_tgz.unlink(missing_ok=True)
+
+ for legacy_name in ("llama-server", "llama-server.tgz"):
+ legacy = drive_cache / legacy_name
+ if legacy.is_file():
+ print(f"[llama] Xóa cache cũ: {legacy_name} ({_fmt_size(legacy.stat().st_size)})")
+ legacy.unlink(missing_ok=True)
+
+ if not allow_build:
+ raise FileNotFoundError(
+ f"Không có llama-server MTP (>= build {MIN_LLAMA_BUILD_FOR_MTP}). "
+ f"Đặt allow_build=True hoặc FORCE_LLAMA_REBUILD=True."
+ )
+
+ bin_path = _build_llama_server(llama_dir, drive_cache, fresh_source=True)
+ setup_ld_library_path(bin_path)
+ _print_bin_info(bin_path, require_mtp=require_mtp)
+ return bin_path
diff --git a/hf-upload/translate_srt.py b/hf-upload/translate_srt.py
new file mode 100644
index 0000000000000000000000000000000000000000..6361e5cbf994f222a6ee371b2bd7b461326d2a67
--- /dev/null
+++ b/hf-upload/translate_srt.py
@@ -0,0 +1,2279 @@
+"""Correct and translate SRT subtitles using Gemma 4 vision + scene context.
+
+Key difference from the old version: a single persistent ``llama-server`` is
+started once (model stays in VRAM) and every cue/scene is handled through HTTP
+requests. The old version spawned ``llama-cli.exe`` per cue, reloading the 6.4 GB
+model every time, which is why a 10-minute video took ~1 hour. With a persistent
+server the same job runs in a few minutes.
+
+Pipeline:
+ parse_srt -> group_scenes -> extract frames (ffmpeg) ->
+ [pass 1] OCR + correct source text per scene (vision, optional) ->
+ [optional] merge adjacent cue pairs (wider timestamps for dubbing) ->
+ [pass 2] scene vision (multi-frame) + relationship lock + translate ->
+ write translated SRT, corrected source SRT, and a JSON report.
+"""
+
+from __future__ import annotations
+
+import argparse
+import atexit
+import base64
+import json
+import re
+import socket
+import subprocess
+import sys
+import tempfile
+import time
+import urllib.error
+import urllib.request
+from collections import Counter
+from dataclasses import dataclass, field
+from difflib import SequenceMatcher
+from pathlib import Path
+from typing import Any
+
+ROOT = Path(__file__).resolve().parent
+MODELS_DIR = ROOT / "models"
+DEFAULT_LLAMA_SERVER = ROOT / "tools" / "llama-cuda" / "llama-server.exe"
+DEFAULT_MODEL = MODELS_DIR / "gemma-4-12B-it-qat-UD-Q4_K_XL.gguf"
+DEFAULT_MMPROJ = MODELS_DIR / "mmproj-F16.gguf"
+DEFAULT_DRAFT = MODELS_DIR / "mtp-gemma-4-12B-it.gguf"
+
+# Hide console windows for child processes on Windows (llama-server, ffmpeg).
+_SUBPROCESS_FLAGS = (
+ subprocess.CREATE_NO_WINDOW if sys.platform == "win32" else 0
+)
+
+SRT_TIME = re.compile(
+ r"(\d{2}):(\d{2}):(\d{2})[,.](\d{3})\s*-->\s*(\d{2}):(\d{2}):(\d{2})[,.](\d{3})"
+)
+
+
+# --------------------------------------------------------------------------- #
+# Data model
+# --------------------------------------------------------------------------- #
+@dataclass
+class Cue:
+ index: int
+ start: float
+ end: float
+ text: str
+ corrected_source: str = ""
+ was_corrected: bool = False
+ ocr_text: str = ""
+ visual_context: str = ""
+ correction_reason: str = ""
+ translated: str = ""
+ honorific_notes: str = ""
+ char_budget: int = 0
+ timing_notes: str = ""
+ speaker: str = ""
+ speaker_gender: str = ""
+ speaker_age_group: str = ""
+ merged_from: list[int] = field(default_factory=list)
+
+ @property
+ def duration(self) -> float:
+ return self.end - self.start
+
+ @property
+ def midpoint(self) -> float:
+ return (self.start + self.end) / 2.0
+
+ @property
+ def start_ts(self) -> str:
+ return _seconds_to_srt_time(self.start)
+
+ @property
+ def end_ts(self) -> str:
+ return _seconds_to_srt_time(self.end)
+
+ @property
+ def source(self) -> str:
+ return self.corrected_source or self.text
+
+
+@dataclass
+class Scene:
+ cues: list[Cue] = field(default_factory=list)
+
+ @property
+ def start(self) -> float:
+ return self.cues[0].start
+
+ @property
+ def end(self) -> float:
+ return self.cues[-1].end
+
+ @property
+ def midpoint(self) -> float:
+ return (self.start + self.end) / 2.0
+
+
+# --------------------------------------------------------------------------- #
+# SRT parsing / writing
+# --------------------------------------------------------------------------- #
+def _srt_time_to_seconds(h: str, m: str, s: str, ms: str) -> float:
+ return int(h) * 3600 + int(m) * 60 + int(s) + int(ms) / 1000.0
+
+
+def _seconds_to_srt_time(seconds: float) -> str:
+ if seconds < 0:
+ seconds = 0.0
+ ms = int(round(seconds * 1000))
+ h, ms = divmod(ms, 3600_000)
+ m, ms = divmod(ms, 60_000)
+ s, ms = divmod(ms, 1000)
+ return f"{h:02d}:{m:02d}:{s:02d},{ms:03d}"
+
+
+def parse_srt(path: Path) -> list[Cue]:
+ raw = path.read_text(encoding="utf-8-sig", errors="replace")
+ blocks = re.split(r"\n\s*\n", raw.strip())
+ cues: list[Cue] = []
+ idx = 0
+ for block in blocks:
+ lines = block.splitlines()
+ m = None
+ time_line = -1
+ for i, ln in enumerate(lines):
+ m = SRT_TIME.search(ln)
+ if m:
+ time_line = i
+ break
+ if not m:
+ continue
+ start = _srt_time_to_seconds(m.group(1), m.group(2), m.group(3), m.group(4))
+ end = _srt_time_to_seconds(m.group(5), m.group(6), m.group(7), m.group(8))
+ text = "\n".join(lines[time_line + 1 :]).strip()
+ idx += 1
+ cues.append(Cue(index=idx, start=start, end=end, text=text))
+ return cues
+
+
+def write_srt(cues: list[Cue], path: Path, use_translation: bool) -> None:
+ out: list[str] = []
+ for i, cue in enumerate(cues, start=1):
+ body = cue.translated if use_translation else cue.source
+ out.append(str(i))
+ out.append(f"{cue.start_ts} --> {cue.end_ts}")
+ out.append(body.strip())
+ out.append("")
+ path.write_text("\n".join(out).strip() + "\n", encoding="utf-8")
+
+
+def write_report(cues: list[Cue], path: Path) -> None:
+ data = [
+ {
+ "index": c.index,
+ "start": c.start_ts,
+ "end": c.end_ts,
+ "original": c.text,
+ "corrected_source": c.corrected_source,
+ "was_corrected": c.was_corrected,
+ "ocr_text": c.ocr_text,
+ "visual_context": c.visual_context,
+ "correction_reason": c.correction_reason,
+ "translated": c.translated,
+ "honorific_notes": c.honorific_notes,
+ "speaker": c.speaker,
+ "speaker_gender": c.speaker_gender,
+ "speaker_age_group": c.speaker_age_group,
+ "duration_sec": round(c.duration, 2),
+ "char_budget": c.char_budget,
+ "translated_chars": len(c.translated.replace("\n", "")),
+ "within_budget": len(c.translated.replace("\n", "")) <= c.char_budget
+ if c.char_budget
+ else True,
+ "timing_notes": c.timing_notes,
+ "merged_from": c.merged_from,
+ }
+ for c in cues
+ ]
+ path.write_text(json.dumps(data, ensure_ascii=False, indent=2), encoding="utf-8")
+
+
+# --------------------------------------------------------------------------- #
+# Scene grouping
+# --------------------------------------------------------------------------- #
+def group_scenes(
+ cues: list[Cue], max_gap: float, max_cues: int, max_duration: float
+) -> list[Scene]:
+ scenes: list[Scene] = []
+ current: list[Cue] = []
+ for cue in cues:
+ if not current:
+ current = [cue]
+ continue
+ gap = cue.start - current[-1].end
+ span = cue.end - current[0].start
+ if gap > max_gap or len(current) >= max_cues or span > max_duration:
+ scenes.append(Scene(current))
+ current = [cue]
+ else:
+ current.append(cue)
+ if current:
+ scenes.append(Scene(current))
+ return scenes
+
+
+def _join_subtitle_text(a: str, b: str) -> str:
+ a, b = a.strip(), b.strip()
+ if not a:
+ return b
+ if not b:
+ return a
+ if a.endswith(("\n",)):
+ return f"{a.rstrip()}\n{b}"
+ if a[-1] in ",。!?、;:…" or b[0] in ",。!?、;:…":
+ return f"{a}{b}"
+ if any("\u4e00" <= ch <= "\u9fff" for ch in (a[-1], b[0])):
+ return f"{a}{b}"
+ return f"{a} {b}"
+
+
+# Characters / particles that mark the END of a spoken sentence in the source.
+# When a cue already ends a sentence, the next cue starts a new thought and the
+# two should not be glued together.
+_SENT_FINAL_PUNCT = "。..!!??…⋯~"
+_SENT_FINAL_PARTICLE = ("吗", "呢", "吧")
+
+
+def _src_text(c: "Cue") -> str:
+ return (c.corrected_source or c.text or "").strip()
+
+
+def _visible_len(s: str) -> int:
+ """Length ignoring whitespace/newlines (CJK chars count as 1)."""
+ return len(re.sub(r"\s+", "", s))
+
+
+def _ends_sentence(s: str) -> bool:
+ s = s.rstrip()
+ if not s:
+ return False
+ if s[-1] in _SENT_FINAL_PUNCT:
+ return True
+ return s.endswith(_SENT_FINAL_PARTICLE)
+
+
+def _speakers_compatible(a: "Cue", b: "Cue") -> bool:
+ """Same speaker, or at least one is unknown (diarization missed it)."""
+ if a.speaker and b.speaker:
+ return a.speaker == b.speaker
+ return True
+
+
+def merge_adjacent_pairs(
+ cues: list[Cue],
+ *,
+ max_gap: float = 0.12,
+ max_combined_duration: float = 10.0,
+ max_combined_chars: int = 32,
+ fragment_chars: int = 12,
+) -> tuple[list[Cue], int]:
+ """Glue a fragmented cue back to its neighbour, but only when it is sensible.
+
+ The source SRT is often split mid-sentence into tiny cues. Merging at most two
+ consecutive cues reconstructs a readable line for dubbing, but only when:
+ - the cues are near-contiguous (next starts where this one ends) and the
+ combined duration fits — a real time gap means a separate utterance;
+ - both cues belong to the same speaker (never merge across speakers);
+ - the first cue does NOT already end a sentence (punctuation / 吗呢吧);
+ - the combined line stays short enough to fit on screen;
+ - at least one side is a short fragment (so two full sentences stay apart).
+ """
+ if len(cues) < 2:
+ return cues, 0
+
+ out: list[Cue] = []
+ merges = 0
+ i = 0
+ while i < len(cues):
+ if i + 1 < len(cues):
+ a, b = cues[i], cues[i + 1]
+ gap = b.start - a.end
+ combined_dur = b.end - a.start
+ src_a = _src_text(a)
+ src_b = _src_text(b)
+ len_a = _visible_len(src_a)
+ len_b = _visible_len(src_b)
+ if (
+ -0.05 <= gap <= max_gap
+ and combined_dur <= max_combined_duration
+ and _speakers_compatible(a, b)
+ and not _ends_sentence(src_a)
+ and (len_a + len_b) <= max_combined_chars
+ and min(len_a, len_b) <= fragment_chars
+ ):
+ merged_text = _join_subtitle_text(src_a, src_b)
+ # Keep the speaker of the longer-spoken cue; note both when they differ.
+ if a.speaker and b.speaker and a.speaker != b.speaker:
+ dominant = a if a.duration >= b.duration else b
+ merged_speaker = dominant.speaker
+ merged_gender = dominant.speaker_gender
+ merged_age = dominant.speaker_age_group
+ else:
+ merged_speaker = a.speaker or b.speaker
+ merged_gender = a.speaker_gender or b.speaker_gender
+ merged_age = a.speaker_age_group or b.speaker_age_group
+ out.append(
+ Cue(
+ index=0,
+ start=a.start,
+ end=b.end,
+ text=_join_subtitle_text(a.text, b.text),
+ corrected_source=merged_text,
+ was_corrected=a.was_corrected or b.was_corrected,
+ speaker=merged_speaker,
+ speaker_gender=merged_gender,
+ speaker_age_group=merged_age,
+ merged_from=[a.index, b.index],
+ )
+ )
+ merges += 1
+ i += 2
+ continue
+ out.append(cues[i])
+ i += 1
+
+ for j, cue in enumerate(out, start=1):
+ cue.index = j
+ return out, merges
+
+
+# --------------------------------------------------------------------------- #
+# ffmpeg frame extraction
+# --------------------------------------------------------------------------- #
+def _run_ffmpeg(args: list[str], timeout: int = 60) -> None:
+ cmd = ["ffmpeg", "-hide_banner", "-loglevel", "error", *args]
+ proc = subprocess.run(
+ cmd, capture_output=True, timeout=timeout, creationflags=_SUBPROCESS_FLAGS
+ )
+ if proc.returncode != 0:
+ raise RuntimeError(
+ "ffmpeg failed: " + proc.stderr.decode("utf-8", "replace").strip()
+ )
+
+
+def extract_frame(
+ video: Path,
+ timestamp: float,
+ output: Path,
+ crop_subtitle: bool,
+ max_width: int = 768,
+) -> Path:
+ vf = []
+ if crop_subtitle:
+ vf.append("crop=iw:ih*0.30:0:ih*0.70")
+ vf.append(f"scale='min({max_width},iw)':-2")
+ args = [
+ "-ss",
+ f"{timestamp:.3f}",
+ "-i",
+ str(video),
+ "-frames:v",
+ "1",
+ "-vf",
+ ",".join(vf),
+ "-q:v",
+ "3",
+ "-y",
+ str(output),
+ ]
+ _run_ffmpeg(args)
+ if not output.exists():
+ raise RuntimeError(f"Failed to extract frame at {timestamp:.3f}s")
+ return output
+
+
+def _image_data_url(path: Path) -> str:
+ data = base64.b64encode(path.read_bytes()).decode("ascii")
+ return f"data:image/jpeg;base64,{data}"
+
+
+# --------------------------------------------------------------------------- #
+# JSON extraction helpers
+# --------------------------------------------------------------------------- #
+_THOUGHT = re.compile(r"<\|channel\|?>thought.*?<\|?channel\|>", re.DOTALL)
+_FENCE = re.compile(r"```(?:json)?\s*([\s\S]*?)\s*```")
+
+
+def _strip_noise(text: str) -> str:
+ text = _THOUGHT.sub("", text)
+ text = re.sub(r"<\|[^>]*\|?>", "", text)
+ return text.strip()
+
+
+def _extract_json(text: str) -> Any:
+ if not text or not text.strip():
+ raise ValueError("Empty model response")
+ cleaned = _strip_noise(text)
+ candidates: list[str] = []
+ fence = _FENCE.search(cleaned)
+ if fence:
+ candidates.append(fence.group(1))
+ candidates.append(cleaned)
+ # Outermost {...} or [...] block. Prefer whichever bracket opens first so an
+ # array that wraps objects is not mistaken for a single inner object.
+ blocks: list[tuple[int, str]] = []
+ for opener, closer in (("{", "}"), ("[", "]")):
+ start = cleaned.find(opener)
+ end = cleaned.rfind(closer)
+ if start != -1 and end != -1 and end > start:
+ blocks.append((start, cleaned[start : end + 1]))
+ for _, block in sorted(blocks, key=lambda b: b[0]):
+ candidates.append(block)
+ for cand in candidates:
+ try:
+ return json.loads(cand)
+ except json.JSONDecodeError:
+ continue
+ raise ValueError("Could not parse JSON from model output:\n" + text)
+
+
+# --------------------------------------------------------------------------- #
+# Persistent llama-server client
+# --------------------------------------------------------------------------- #
+def _free_port() -> int:
+ with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
+ s.bind(("127.0.0.1", 0))
+ return s.getsockname()[1]
+
+
+class LlamaServer:
+ """Starts (or reuses) a persistent llama-server and talks to it over HTTP."""
+
+ def __init__(
+ self,
+ server_bin: Path,
+ model: Path,
+ mmproj: Path,
+ ngl: int = 999,
+ ctx_size: int = 8192,
+ port: int | None = None,
+ draft: Path | None = None,
+ use_mtp: bool = True,
+ temperature: float = 0.7,
+ max_tokens: int = 1024,
+ start_timeout: int = 180,
+ mtp_start_timeout: int = 420,
+ existing_url: str | None = None,
+ ) -> None:
+ self.server_bin = server_bin
+ self.model = model
+ self.mmproj = mmproj
+ self.ngl = ngl
+ self.ctx_size = ctx_size
+ self.draft = draft
+ self.use_mtp = use_mtp
+ self.temperature = temperature
+ self.max_tokens = max_tokens
+ self.start_timeout = start_timeout
+ self.mtp_start_timeout = mtp_start_timeout
+ self.proc: subprocess.Popen | None = None
+
+ if existing_url:
+ self.base_url = existing_url.rstrip("/")
+ self.owns_server = False
+ else:
+ self.port = port or _free_port()
+ self.base_url = f"http://127.0.0.1:{self.port}"
+ self.owns_server = True
+
+ # -- lifecycle -------------------------------------------------------- #
+ def _server_args(self, with_mtp: bool) -> list[str]:
+ args = [
+ str(self.server_bin),
+ "-m",
+ str(self.model),
+ "--mmproj",
+ str(self.mmproj),
+ "-ngl",
+ str(self.ngl),
+ "-fa",
+ "on",
+ "-c",
+ str(self.ctx_size),
+ "--host",
+ "127.0.0.1",
+ "--port",
+ str(self.port),
+ "--parallel",
+ "1",
+ "--reasoning",
+ "off",
+ "--reasoning-budget",
+ "0",
+ ]
+ if with_mtp and self.draft and self.draft.exists():
+ args += [
+ "--model-draft",
+ str(self.draft),
+ "--spec-type",
+ "draft-mtp",
+ "--spec-draft-n-max",
+ "4",
+ "--spec-draft-ngl",
+ str(self.ngl),
+ "-ngld",
+ str(self.ngl),
+ ]
+ return args
+
+ def _wait_healthy(self, timeout: int) -> bool:
+ deadline = time.time() + timeout
+ url = f"{self.base_url}/health"
+ while time.time() < deadline:
+ if self.proc is not None and self.proc.poll() is not None:
+ return False
+ try:
+ with urllib.request.urlopen(url, timeout=3) as resp:
+ if resp.status == 200:
+ return True
+ except (urllib.error.URLError, OSError):
+ time.sleep(1.0)
+ return False
+
+ def start(self) -> None:
+ if not self.owns_server:
+ if not self._wait_healthy(15):
+ raise RuntimeError(f"No healthy server at {self.base_url}")
+ return
+
+ attempts = [True, False] if self.use_mtp else [False]
+ last_log = ""
+ for with_mtp in attempts:
+ label = "with MTP" if with_mtp else "without MTP"
+ print(f"[server] starting llama-server ({label})...", flush=True)
+ log = tempfile.NamedTemporaryFile(
+ prefix="llama_server_", suffix=".log", delete=False, mode="w"
+ )
+ self._log_path = Path(log.name)
+ self.proc = subprocess.Popen(
+ self._server_args(with_mtp),
+ stdout=log,
+ stderr=subprocess.STDOUT,
+ cwd=str(self.server_bin.parent),
+ creationflags=_SUBPROCESS_FLAGS,
+ )
+ atexit.register(self.stop)
+ timeout = self.mtp_start_timeout if with_mtp else self.start_timeout
+ if self._wait_healthy(timeout):
+ print(f"[server] ready at {self.base_url} ({label})", flush=True)
+ return
+ self.stop()
+ last_log = self._log_path.read_text(errors="replace")[-2000:]
+ print(f"[server] failed to start {label}; retrying...", flush=True)
+ if last_log.strip():
+ print("[server] log (tail):\n" + last_log, flush=True)
+ raise RuntimeError("llama-server did not become healthy.\n" + last_log)
+
+ def stop(self) -> None:
+ if self.proc is not None and self.proc.poll() is None:
+ self.proc.terminate()
+ try:
+ self.proc.wait(timeout=15)
+ except subprocess.TimeoutExpired:
+ self.proc.kill()
+ self.proc = None
+
+ # -- requests --------------------------------------------------------- #
+ def chat(
+ self,
+ prompt: str,
+ images: list[Path] | None = None,
+ system: str | None = None,
+ max_tokens: int | None = None,
+ retries: int = 2,
+ temperature: float | None = None,
+ ) -> str:
+ content: list[dict[str, Any]] = []
+ for img in images or []:
+ content.append(
+ {"type": "image_url", "image_url": {"url": _image_data_url(img)}}
+ )
+ content.append({"type": "text", "text": prompt})
+
+ messages: list[dict[str, Any]] = []
+ if system:
+ messages.append({"role": "system", "content": system})
+ messages.append({"role": "user", "content": content})
+
+ payload = {
+ "messages": messages,
+ "temperature": self.temperature if temperature is None else temperature,
+ "top_p": 0.95,
+ "top_k": 64,
+ "max_tokens": max_tokens or self.max_tokens,
+ "cache_prompt": True,
+ "stream": False,
+ }
+ body = json.dumps(payload).encode("utf-8")
+ url = f"{self.base_url}/v1/chat/completions"
+
+ last_err: Exception | None = None
+ for attempt in range(retries + 1):
+ try:
+ req = urllib.request.Request(
+ url, data=body, headers={"Content-Type": "application/json"}
+ )
+ with urllib.request.urlopen(req, timeout=600) as resp:
+ data = json.loads(resp.read().decode("utf-8"))
+ msg = data["choices"][0]["message"]
+ content = msg.get("content") or ""
+ if not content.strip():
+ # Some templates route everything into reasoning_content.
+ content = msg.get("reasoning_content") or ""
+ if not content.strip():
+ raise ValueError("empty content")
+ return content
+ except (urllib.error.URLError, OSError, KeyError, ValueError) as exc:
+ last_err = exc
+ time.sleep(2.0 * (attempt + 1))
+ raise RuntimeError(f"Chat request failed: {last_err}")
+
+ def chat_json(
+ self,
+ prompt: str,
+ images: list[Path] | None = None,
+ system: str | None = None,
+ max_tokens: int | None = None,
+ temperature: float | None = None,
+ ) -> Any:
+ text = self.chat(
+ prompt,
+ images=images,
+ system=system,
+ max_tokens=max_tokens,
+ temperature=temperature,
+ )
+ return _extract_json(text)
+
+
+# --------------------------------------------------------------------------- #
+# Subtitle timing helpers
+# --------------------------------------------------------------------------- #
+def subtitle_char_budget(
+ duration: float,
+ chars_per_sec: float = 17.0,
+ max_line_chars: int = 42,
+ max_lines: int = 2,
+ min_budget: int = 6,
+) -> int:
+ """Estimate how many characters fit comfortably on screen for *duration* seconds."""
+ if duration <= 0:
+ return min_budget
+ cps_budget = int(duration * chars_per_sec)
+ display_cap = max_line_chars * max_lines
+ return max(min_budget, min(cps_budget, display_cap))
+
+
+def _translated_char_count(text: str) -> int:
+ return len(text.replace("\n", ""))
+
+
+def _normalize_subtitle_text(text: str) -> str:
+ """Single-line subtitles: collapse model-inserted line breaks."""
+ return re.sub(r"\s*\n\s*", " ", text).strip()
+
+
+def _apply_char_budgets(
+ cues: list[Cue], chars_per_sec: float, max_line_chars: int, max_lines: int
+) -> None:
+ for c in cues:
+ c.char_budget = subtitle_char_budget(
+ c.duration, chars_per_sec, max_line_chars, max_lines
+ )
+
+
+# --------------------------------------------------------------------------- #
+# Prompts
+# --------------------------------------------------------------------------- #
+CORRECTION_SYSTEM = (
+ "You are a meticulous subtitle proofreader. You look at video frames that may "
+ "contain burned-in (hard) subtitles and compare them with the provided SRT text. "
+ "Reply with JSON only."
+)
+
+TRANSLATION_SYSTEM = (
+ "You are an expert Vietnamese film/TV subtitle translator. Write the way native "
+ "Vietnamese speakers actually talk: natural, idiomatic, smooth and CONCISE spoken "
+ "language — never word-for-word, stiff, or translationese. Convey the full meaning, "
+ "tone, emotion and humor, but cut redundant words, drop subjects/pronouns that are "
+ "obvious from context, and avoid clunky calques (e.g. 'thông qua việc', 'một cách', "
+ "repeating 'tôi ... tôi ...'). A shorter, natural line is better than a long, "
+ "faithful-but-awkward one. Keep correct forms of address / pronouns. Reply with JSON only."
+)
+
+
+def build_correction_prompt(scene: Scene, source_lang: str) -> str:
+ lines = []
+ for i, c in enumerate(scene.cues):
+ lines.append(
+ f"- image_index={i} | cue_index={c.index} | "
+ f"{c.start_ts} --> {c.end_ts} | srt_text: {c.text!r}"
+ )
+ src = "the same language" if source_lang == "auto" else source_lang
+ return (
+ f"Source language: {src}.\n"
+ "Each cue below has a corresponding cropped video frame (in the same order as "
+ "the images). Read any burned-in/hard subtitle text visible in each frame and "
+ "compare it with the SRT text. If the SRT text is wrong (typos, missing words, "
+ "misrecognized characters), fix it to match what is actually shown/spoken. If "
+ "the frame has no visible subtitle, keep the SRT text unchanged.\n\n"
+ "Cues:\n" + "\n".join(lines) + "\n\n"
+ 'Return ONLY a JSON array, one object per cue, in this exact shape:\n'
+ '[{"cue_index": , "corrected_source": "", '
+ '"was_corrected": , "ocr_text": "", '
+ '"visual_context": "", '
+ '"correction_reason": ""}]'
+ )
+
+
+_GENDER_LABEL = {
+ "male": "nam",
+ "female": "nữ",
+ "child": "trẻ em",
+ "unknown": "chưa rõ",
+}
+
+_AGE_GROUP_LABEL = {
+ "child": "trẻ em",
+ "teen": "thiếu niên",
+ "young_adult": "thanh niên",
+ "middle_aged": "trung niên",
+ "elderly": "lớn tuổi",
+}
+
+
+def _gender_label(gender: str) -> str:
+ return _GENDER_LABEL.get((gender or "").lower(), "chưa rõ")
+
+
+def _age_label(age_group: str) -> str:
+ return _AGE_GROUP_LABEL.get((age_group or "").lower(), "")
+
+
+def _speaker_descriptor(gender: str, age_group: str) -> str:
+ """e.g. 'nữ, thanh niên' or 'nam' or 'chưa rõ'."""
+ parts = [_gender_label(gender)]
+ age = _age_label(age_group)
+ if age:
+ parts.append(age)
+ return ", ".join(parts)
+
+
+def _speaker_tag(cue: Cue) -> str:
+ """Compact 'who is speaking' tag, e.g. 'SPEAKER_01 (nữ, thanh niên)'."""
+ if not cue.speaker:
+ return ""
+ return f"{cue.speaker} ({_speaker_descriptor(cue.speaker_gender, cue.speaker_age_group)})"
+
+
+def build_speaker_registry(cues: list[Cue]) -> str:
+ """List every detected speaker + gender + age so pronouns stay stable."""
+ seen: dict[str, tuple[str, str]] = {}
+ for c in cues:
+ if c.speaker and c.speaker not in seen:
+ seen[c.speaker] = (c.speaker_gender, c.speaker_age_group)
+ if not seen:
+ return ""
+ rows = [
+ f"- {spk}: {_speaker_descriptor(gender, age_group)}"
+ for spk, (gender, age_group) in seen.items()
+ ]
+ return "\n".join(rows)
+
+
+# --------------------------------------------------------------------------- #
+# Character name consistency (CJK source -> locked Vietnamese spelling)
+# --------------------------------------------------------------------------- #
+_CJK_NAME_TOKEN = re.compile(r"[\u4e00-\u9fff]{2,3}")
+_CJK_STOPWORDS = frozenset(
+ {
+ "什么", "怎么", "没有", "不是", "这个", "那个", "自己", "时候", "知道",
+ "可以", "已经", "因为", "所以", "但是", "如果", "我们", "你们", "他们",
+ "这样", "为什么", "一下", "真的", "还是", "就是", "不要", "现在", "以前",
+ "今天", "明天", "一起", "东西", "事情", "地方", "个人", "大家", "对不起",
+ "谢谢", "没关系", "等等", "不对", "当然", "可能", "应该", "一定", "起来",
+ "出去", "回来", "告诉", "觉得", "认为", "喜欢", "希望", "开始", "结束",
+ "是不是", "好不好", "能不能", "会不会", "干什么", "怎么办", "怎么样",
+ "在吗", "是吗", "对吧", "行了", "算了", "没事", "好了",
+ # Common nouns/verbs/adverbs that look like 2-3 char names but are not.
+ "别墅", "礼物", "到了", "收尾", "这里", "那里", "这边", "那边", "哪里",
+ "一个", "两个", "怎样", "多少", "一点", "有点", "一些", "这些", "那些",
+ "时间", "问题", "工作", "公司", "医院", "手术", "电话", "老师", "医生",
+ "护士", "患者", "病人", "孩子", "生日", "结婚", "离婚", "生活", "感觉",
+ "关系", "办法", "方法", "样子", "意思", "消息", "情况", "决定", "准备",
+ "需要", "然后", "还有", "而且", "不过", "只是", "其实", "当时", "后来",
+ "之后", "之前", "一直", "马上", "立刻", "突然", "终于", "总是", "经常",
+ "也许", "大概", "也是", "刚才", "刚刚", "最后", "最近", "以后", "以为",
+ "妈妈", "爸爸", "母亲", "父亲", "儿子", "女儿", "老公", "老婆", "妻子",
+ "丈夫", "先生", "太太", "小姐", "样东", "东西", "时候", "事情", "钱财",
+ "资金", "公里", "回家", "出来", "进来", "下来", "上来", "过来", "过去",
+ "不会", "不能", "不想", "不用", "明白", "清楚", "记得", "忘记", "答应",
+ }
+)
+_VN_NAME_WORD = (
+ r"[A-ZÀÁẢÃẠĂẮẰẲẴẶÂẤẦẨẪẬĐÉÈẺẼẸÊẾỀỂỄỆÍÌỈĨỊÓÒỎÕỌÔỐỒỔỖỘƠỚỜỞỠỢÚÙỦŨỤƯỨỪỬỮỰÝỲỶỸỴ]"
+ r"[a-zàáảãạăắằẳẵặâấầẩẫậđéèẻẽẹêếềểễệíìỉĩịóòỏõọôốồổỗộơớờởỡợúùủũụưứừửữựýỳỷỹỵ]*"
+)
+_VN_NAME_RE = re.compile(rf"(?:{_VN_NAME_WORD})(?:\s+(?:{_VN_NAME_WORD})){{1,3}}")
+_VN_NAME_SKIP_FIRST = frozenset(
+ {
+ # Kinship / personal pronouns that precede names but are not part of them.
+ "Anh", "Em", "Cô", "Chú", "Bác", "Ông", "Bà", "Chị", "Dì", "Cậu",
+ "Mợ", "Thím", "Cụ", "Bố", "Ba", "Mẹ", "Má", "Con", "Cháu",
+ "Tôi", "Mình", "Tao", "Mày", "Ta", "Họ", "Nó", "Hắn", "Y", "Chúng",
+ # Sentence-initial connectives / adverbs that get Title-cased before a
+ # name (e.g. "Nên Cố Hàn Thâm" = "Cho nên, Cố Hàn Thâm...").
+ "Nhưng", "Vậy", "Được", "Không", "Có", "Rồi", "Sao", "Vì", "Nếu",
+ "Khi", "Mà", "Thì", "Hay", "Hoặc", "Nên", "Cũng", "Đã", "Đang", "Sẽ",
+ "Lại", "Vẫn", "Còn", "Chỉ", "Cứ", "Phải", "Rất", "Quá", "Thật",
+ "Chính", "Của", "Cho", "Với", "Và", "Trong", "Này", "Đó", "Kia", "Ai",
+ "Gì", "Đây", "Thế", "Tại", "Bởi", "Do", "Theo", "Từ", "Đến", "Về",
+ "Lúc", "Giờ", "Mai", "Nay", "Nãy", "Sau", "Trước", "Trên", "Dưới",
+ "Ngoài", "Bên", "Cùng", "Mọi", "Mỗi", "Cả", "Những", "Các", "Một",
+ "Hai", "Bốn", "Năm", "Người", "Việc", "Chuyện", "Để", "Là", "Khiến",
+ }
+)
+_NAME_SIMILARITY_THRESHOLD = 0.68
+
+
+def _extract_cjk_name_tokens(text: str) -> list[str]:
+ return [
+ tok
+ for tok in _CJK_NAME_TOKEN.findall(text or "")
+ if tok not in _CJK_STOPWORDS
+ ]
+
+
+def _extract_vn_name_phrases(text: str) -> list[str]:
+ names: list[str] = []
+ for m in _VN_NAME_RE.finditer(text or ""):
+ parts = m.group(0).strip().split()
+ # Strip leading kinship/pronoun/connective words that get Title-cased in
+ # front of a name (e.g. "Nên Cố Hàn Thâm" -> "Cố Hàn Thâm", a sentence
+ # that begins with "Nên" = "cho nên"). Repeat to peel multiple, e.g.
+ # "Và Anh Tú" -> "Tú" would over-strip, so stop once a non-skip word is hit.
+ while parts and parts[0] in _VN_NAME_SKIP_FIRST:
+ parts.pop(0)
+ if len(parts) < 2:
+ continue
+ names.append(" ".join(parts))
+ return names
+
+
+def _name_part_similarity(a: str, b: str) -> float:
+ return SequenceMatcher(None, a.lower(), b.lower()).ratio()
+
+
+def _name_phrase_similarity(a: str, b: str) -> float:
+ a_parts = a.split()
+ b_parts = b.split()
+ if not a_parts or not b_parts:
+ return 0.0
+ if len(a_parts) == len(b_parts):
+ scores = [_name_part_similarity(x, y) for x, y in zip(a_parts, b_parts)]
+ return sum(scores) / len(scores)
+ if abs(len(a_parts) - len(b_parts)) != 1:
+ return 0.0
+ short, long = (
+ (a_parts, b_parts) if len(a_parts) < len(b_parts) else (b_parts, a_parts)
+ )
+ best = 0.0
+ for i in range(len(long) - len(short) + 1):
+ chunk = long[i : i + len(short)]
+ scores = [_name_part_similarity(x, y) for x, y in zip(short, chunk)]
+ best = max(best, sum(scores) / len(scores))
+ return best
+
+
+@dataclass
+class NameRegistry:
+ """Lock one Vietnamese spelling per source proper name (CJK in SRT)."""
+
+ source_tokens: set[str] = field(default_factory=set)
+ source_to_vn: dict[str, str] = field(default_factory=dict)
+ aliases: dict[str, str] = field(default_factory=dict)
+
+ @classmethod
+ def from_cues(cls, cues: list[Cue]) -> "NameRegistry":
+ counts: Counter[str] = Counter()
+ for c in cues:
+ counts.update(_extract_cjk_name_tokens(c.source))
+ # Only tokens that RECUR are treated as proper names. The old rule that
+ # also kept any 2-char token seen once flooded the registry with common
+ # words (别墅=villa, 礼物=gift, 到了=arrived...) that then corrupted real
+ # names. Genuine character/place names recur across a film.
+ tokens = {tok for tok, n in counts.items() if n >= 2}
+ return cls(source_tokens=tokens)
+
+ @property
+ def has_names(self) -> bool:
+ return bool(self.source_tokens or self.source_to_vn)
+
+ def _register(self, source: str, canonical: str) -> None:
+ canonical = canonical.strip()
+ if not source or not canonical:
+ return
+ prev = self.source_to_vn.get(source)
+ if prev and prev != canonical:
+ self._add_alias(canonical, prev)
+ return
+ self.source_to_vn[source] = canonical
+
+ def _add_alias(self, wrong: str, canonical: str) -> None:
+ wrong = wrong.strip()
+ canonical = canonical.strip()
+ if not wrong or wrong == canonical:
+ return
+ self.aliases[wrong] = canonical
+
+ def _pick_canonical_for_source(self, vn_names: list[str]) -> str:
+ if not vn_names:
+ return ""
+ scored: list[tuple[float, str]] = []
+ for name in vn_names:
+ score = float(len(name.split()))
+ for other in self.source_to_vn.values():
+ score += _name_phrase_similarity(name, other) * 2.0
+ scored.append((score, name))
+ scored.sort(key=lambda p: (-p[0], -len(p[1])))
+ return scored[0][1]
+
+ def observe_cue(self, cue: Cue) -> None:
+ """Learn mappings from a translated cue and fix drift in-place."""
+ if not cue.translated:
+ return
+ src_names = [
+ tok
+ for tok in _extract_cjk_name_tokens(cue.source)
+ if tok in self.source_tokens or tok in self.source_to_vn
+ ]
+ vn_names = _extract_vn_name_phrases(cue.translated)
+ for src in src_names:
+ if src in self.source_to_vn:
+ canonical = self.source_to_vn[src]
+ for vn in vn_names:
+ if vn != canonical and _name_phrase_similarity(vn, canonical) >= _NAME_SIMILARITY_THRESHOLD:
+ self._add_alias(vn, canonical)
+ elif vn_names:
+ self._register(src, self._pick_canonical_for_source(vn_names))
+ cue.translated = self.normalize_text(cue.translated)
+
+ def normalize_text(self, text: str) -> str:
+ if not text:
+ return text
+ out = text
+ for wrong in sorted(self.aliases, key=len, reverse=True):
+ if wrong in out:
+ out = out.replace(wrong, self.aliases[wrong])
+ for canonical in self.source_to_vn.values():
+ for candidate in _extract_vn_name_phrases(out):
+ if candidate == canonical:
+ continue
+ if _name_phrase_similarity(candidate, canonical) >= _NAME_SIMILARITY_THRESHOLD:
+ out = out.replace(candidate, canonical)
+ return out
+
+ def reconcile_all(self, cues: list[Cue]) -> None:
+ """Final pass: cluster similar Vietnamese names and unify spellings."""
+ phrases: Counter[str] = Counter()
+ for c in cues:
+ phrases.update(_extract_vn_name_phrases(c.translated))
+ canonicals = list(self.source_to_vn.values())
+ ordered = sorted(phrases, key=lambda p: (-phrases[p], -len(p)))
+ for phrase in ordered:
+ if phrase in canonicals:
+ continue
+ best_canon = ""
+ best_score = _NAME_SIMILARITY_THRESHOLD
+ for canon in canonicals:
+ score = _name_phrase_similarity(phrase, canon)
+ if score > best_score:
+ best_score = score
+ best_canon = canon
+ if best_canon:
+ self._add_alias(phrase, best_canon)
+ for c in cues:
+ if c.translated:
+ c.translated = self.normalize_text(c.translated)
+
+ def build_prompt_block(self) -> str:
+ if self.source_to_vn:
+ rows = [
+ f"- {src} → {vn}" for src, vn in sorted(self.source_to_vn.items())
+ ]
+ body = (
+ "\nNAME GUIDE (locked character/place names — use EXACT spellings):\n"
+ + "\n".join(rows)
+ )
+ elif self.source_tokens:
+ sample = ", ".join(sorted(self.source_tokens)[:16])
+ body = (
+ "\nNAME CONSISTENCY (proper names detected in source SRT):\n"
+ f"- Examples in this file: {sample}\n"
+ "- Pick ONE Vietnamese spelling per name and reuse it in every cue."
+ )
+ else:
+ return ""
+ return (
+ body
+ + "\n- When the source mentions a listed name, output ONLY the locked "
+ "Vietnamese form on the right.\n"
+ "- Never alternate spellings for the same person (e.g. Ôn vs Ô, "
+ "Diễn vs Ngiễn, Thời vs Thi, extra/missing middle syllables).\n"
+ "- For a NEW name not yet listed: choose ONE natural Vietnamese "
+ "transliteration and reuse it in every later cue.\n"
+ )
+
+
+# --------------------------------------------------------------------------- #
+# Global forms-of-address resolution (xưng hô) per speaker pair
+# --------------------------------------------------------------------------- #
+# These CJK patterns are NOT used to force a relationship anymore; they only
+# surface kinship / romance keywords as *hints* for the LLM resolver below.
+_CJK_ROMANCE_RE = re.compile(
+ r"同岁|同龄|男朋友|女朋友|男友|女友|老公|老婆|妻子|丈夫|"
+ r"亲爱的|宝贝|宝宝|喜欢你|爱你|我爱你|结婚|订婚|"
+ r"供我上|供你上|养你读|养我读|供我读|供你读|"
+ r"我们(?:俩|两个|一起)|咱俩|情侣|对象"
+)
+_CJK_FAMILY_RE = re.compile(
+ r"妈妈|母亲|妈|爸爸|父亲|爸|爹|娘|"
+ r"儿子|女儿|闺女|孩儿|乖儿子|乖女儿|"
+ r"咱妈|咱爸|老妈|老爸|哥哥|姐姐|弟弟|妹妹|"
+ r"爷爷|奶奶|外婆|外公|姥姥|姥爷|叔叔|阿姨|舅舅|姑姑"
+)
+def _speaker_pair_key(a: str, b: str) -> frozenset[str]:
+ return frozenset({a, b})
+
+
+def _collect_relationship_evidence(texts: list[str]) -> str:
+ """List kinship / romance CJK keywords found in a pair's dialogue (hints only)."""
+ joined = "\n".join(t for t in texts if t)
+
+ def _clean(pattern: re.Pattern[str]) -> list[str]:
+ seen: list[str] = []
+ for m in pattern.finditer(joined):
+ tok = re.sub(r"[^\u4e00-\u9fff]", "", m.group(0))
+ if tok and tok not in seen:
+ seen.append(tok)
+ return seen
+
+ parts: list[str] = []
+ fam = _clean(_CJK_FAMILY_RE)
+ rom = _clean(_CJK_ROMANCE_RE)
+ if fam:
+ parts.append("kinship words: " + ", ".join(fam[:8]))
+ if rom:
+ parts.append("romance/same-age words: " + ", ".join(rom[:8]))
+ return "; ".join(parts)
+
+
+def _conversing_pairs(
+ cues: list[Cue], *, min_exchanges: int = 2
+) -> list[frozenset[str]]:
+ """Pairs of speakers that actually take turns talking to each other."""
+ counts: Counter[frozenset[str]] = Counter()
+ seq = [c.speaker for c in cues if c.speaker]
+ for a, b in zip(seq, seq[1:]):
+ if a and b and a != b:
+ counts[_speaker_pair_key(a, b)] += 1
+ return [pair for pair, n in counts.items() if n >= min_exchanges]
+
+
+def _pair_dialogue_sample(
+ cues: list[Cue], pair: frozenset[str], *, max_lines: int = 40
+) -> list[str]:
+ """Chronological, evenly-spaced dialogue lines for a speaker pair."""
+ lines = [
+ f"[{c.speaker}] {c.source.strip()}"
+ for c in cues
+ if c.speaker in pair and c.source.strip()
+ ]
+ if len(lines) <= max_lines:
+ return lines
+ step = len(lines) / max_lines
+ return [lines[int(i * step)] for i in range(max_lines)]
+
+
+@dataclass
+class AddressLink:
+ """How one speaker refers to themselves and addresses the other."""
+
+ self_term: str = ""
+ other_term: str = ""
+
+
+@dataclass
+class AddressRegistry:
+ """Locked, globally-resolved forms of address per directed speaker pair."""
+
+ directed: dict[tuple[str, str], AddressLink] = field(default_factory=dict)
+ relationships: dict[frozenset[str], dict[str, str]] = field(default_factory=dict)
+
+ @property
+ def has_map(self) -> bool:
+ return bool(self.directed)
+
+ def set_pair(
+ self,
+ a: str,
+ b: str,
+ *,
+ a_self: str,
+ a_other: str,
+ b_self: str,
+ b_other: str,
+ relationship: str = "",
+ confidence: str = "",
+ evidence: str = "",
+ ) -> None:
+ if not a or not b or a == b:
+ return
+ # Need at least the "other" terms (how each addresses the other).
+ if not (a_other.strip() or b_other.strip()):
+ return
+ self.directed[(a, b)] = AddressLink(a_self.strip(), a_other.strip())
+ self.directed[(b, a)] = AddressLink(b_self.strip(), b_other.strip())
+ self.relationships[_speaker_pair_key(a, b)] = {
+ "a": a,
+ "b": b,
+ "relationship": relationship.strip(),
+ "confidence": confidence.strip(),
+ "evidence": evidence.strip(),
+ }
+
+ def describe_lines(self) -> list[str]:
+ out: list[str] = []
+ for meta in self.relationships.values():
+ a, b = meta["a"], meta["b"]
+ ab = self.directed.get((a, b))
+ ba = self.directed.get((b, a))
+ if not ab or not ba:
+ continue
+ rel = meta.get("relationship") or "?"
+ out.append(
+ f"{a}->{b}: {ab.self_term or '?'}/{ab.other_term or '?'} | "
+ f"{b}->{a}: {ba.self_term or '?'}/{ba.other_term or '?'} ({rel})"
+ )
+ return out
+
+ def build_prompt_block(self, scene: Scene | None = None) -> str:
+ if not self.directed:
+ return ""
+ scene_spks = (
+ {c.speaker for c in scene.cues if c.speaker} if scene is not None else None
+ )
+ rows: list[str] = []
+ for meta in self.relationships.values():
+ a, b = meta["a"], meta["b"]
+ if scene_spks is not None and a not in scene_spks and b not in scene_spks:
+ continue
+ ab = self.directed.get((a, b))
+ ba = self.directed.get((b, a))
+ if not ab or not ba:
+ continue
+ rel = meta.get("relationship") or "?"
+ rows.append(
+ f"- {a} <-> {b} ({rel}): when {a} speaks to {b}, {a} calls self "
+ f"\"{ab.self_term}\" and addresses {b} as \"{ab.other_term}\"; "
+ f"when {b} speaks to {a}, {b} calls self \"{ba.self_term}\" and "
+ f"addresses {a} as \"{ba.other_term}\"."
+ )
+ if not rows:
+ return ""
+ return (
+ "\nADDRESS MAP (forms of address locked for the WHOLE film — MANDATORY):\n"
+ + "\n".join(rows)
+ + "\n- Use EXACTLY these self-term / other-term for each speaker pair in "
+ "every cue. NEVER switch to a different pronoun pair between cues or "
+ "mid-scene (no flipping anh/em <-> mẹ/con <-> chị/em <-> cô/cháu...).\n"
+ "- These were resolved from the whole conversation; trust them over voice "
+ "age hints.\n"
+ "- If a cue's speaker pair is NOT listed here, infer the natural Vietnamese "
+ "forms of address from dialogue, names and scene, then keep them stable.\n"
+ )
+
+
+ADDRESS_SYSTEM = (
+ "You are an expert Vietnamese translator specializing in forms of address "
+ "(xưng hô). Given a whole conversation between two characters, decide the "
+ "natural Vietnamese pronoun pair they use, consistently for the entire film. "
+ "Reply with JSON only."
+)
+
+
+def build_address_prompt(
+ spk_a: str,
+ desc_a: str,
+ spk_b: str,
+ desc_b: str,
+ dialogue: list[str],
+ evidence: str,
+ source_lang: str,
+) -> str:
+ src = "the source language" if source_lang == "auto" else source_lang
+ ev = f"\nDetected source keywords (hints, may be wrong): {evidence}\n" if evidence else ""
+ return (
+ f"Two characters talk to each other in {src}. Decide the correct Vietnamese "
+ "forms of address between them and keep them STABLE for the whole film.\n\n"
+ f"Speaker A = {spk_a} (voice guess: {desc_a}).\n"
+ f"Speaker B = {spk_b} (voice guess: {desc_b}).\n"
+ f"{ev}\n"
+ "Conversation sample (chronological):\n"
+ + "\n".join(dialogue)
+ + "\n\nGuidance:\n"
+ "- Infer the relationship from what they SAY (kinship terms like 妈/儿子, "
+ "romance like 老公/同岁, formality), NOT from the voice age guess which is "
+ "often wrong.\n"
+ "- Vietnamese has many pronoun pairs; pick whichever fits naturally: "
+ "vợ/chồng, anh-em, chị-em, mẹ-con, bố-con, ông/bà-cháu, cô/chú/dì/cậu-cháu, "
+ "thầy/cô-em (teacher/student), bạn-tớ/cậu, mày-tao, anh/chị-tôi, etc.\n"
+ "- For EACH direction give how the speaker refers to THEMSELVES (self) and how "
+ "they ADDRESS the other (other), e.g. self=\"anh\", other=\"em\".\n\n"
+ "Return ONLY JSON in this exact shape:\n"
+ '{"relationship": "",'
+ ' "confidence": "high|medium|low",'
+ ' "a_to_b": {"self": "", "other": ""},'
+ ' "b_to_a": {"self": "", "other": ""},'
+ ' "notes": ""}'
+ )
+
+
+def resolve_address_map(
+ client: "LlamaServer",
+ cues: list[Cue],
+ source_lang: str,
+ *,
+ min_exchanges: int = 2,
+ temperature: float = 0.2,
+ log: Any = None,
+) -> AddressRegistry:
+ """Resolve a stable Vietnamese forms-of-address map for each conversing pair."""
+ def _log(msg: str) -> None:
+ if log:
+ log(msg)
+
+ reg = AddressRegistry()
+ genders: dict[str, str] = {}
+ ages: dict[str, str] = {}
+ for c in cues:
+ if c.speaker:
+ genders.setdefault(c.speaker, c.speaker_gender or "")
+ ages.setdefault(c.speaker, c.speaker_age_group or "")
+
+ pairs = _conversing_pairs(cues, min_exchanges=min_exchanges)
+ if not pairs:
+ return reg
+ _log(f" Resolving forms of address for {len(pairs)} speaker pair(s)...")
+ for pair in pairs:
+ a, b = sorted(pair)
+ sample = _pair_dialogue_sample(cues, pair)
+ if not sample:
+ continue
+ evidence = _collect_relationship_evidence(
+ [c.source for c in cues if c.speaker in pair]
+ )
+ prompt = build_address_prompt(
+ a,
+ _speaker_descriptor(genders.get(a, ""), ages.get(a, "")),
+ b,
+ _speaker_descriptor(genders.get(b, ""), ages.get(b, "")),
+ sample,
+ evidence,
+ source_lang,
+ )
+ try:
+ result = client.chat_json(
+ prompt,
+ system=ADDRESS_SYSTEM,
+ max_tokens=512,
+ temperature=temperature,
+ )
+ except Exception as exc: # noqa: BLE001
+ _log(f" [warn] address resolution failed for {a}<->{b}: {exc}")
+ continue
+ if not isinstance(result, dict):
+ continue
+ ab = result.get("a_to_b") if isinstance(result.get("a_to_b"), dict) else {}
+ ba = result.get("b_to_a") if isinstance(result.get("b_to_a"), dict) else {}
+ reg.set_pair(
+ a,
+ b,
+ a_self=str(ab.get("self") or ""),
+ a_other=str(ab.get("other") or ""),
+ b_self=str(ba.get("self") or ""),
+ b_other=str(ba.get("other") or ""),
+ relationship=str(result.get("relationship") or ""),
+ confidence=str(result.get("confidence") or ""),
+ evidence=evidence,
+ )
+ return reg
+
+
+# --------------------------------------------------------------------------- #
+# Pass 2 scene vision (multi-frame)
+# --------------------------------------------------------------------------- #
+SCENE_VISION_SYSTEM = (
+ "You analyze video frames to help subtitle translators. "
+ "Describe who is visible and their apparent relationship. Reply with JSON only."
+)
+
+
+def extract_scene_frames(
+ video: Path,
+ scene: Scene,
+ frame_dir: Path,
+ n_frames: int = 3,
+ max_width: int = 768,
+) -> list[Path]:
+ """Extract n evenly-spaced full frames across a scene (start/mid/end)."""
+ n = max(1, min(n_frames, 5))
+ span = max(scene.end - scene.start, 0.05)
+ if n == 1:
+ timestamps = [scene.midpoint]
+ else:
+ pad = min(0.15, span * 0.08)
+ lo, hi = scene.start + pad, scene.end - pad
+ if hi <= lo:
+ lo, hi = scene.start, scene.end
+ step = (hi - lo) / (n - 1)
+ timestamps = [lo + step * i for i in range(n)]
+ paths: list[Path] = []
+ base = int(scene.start * 1000)
+ for i, ts in enumerate(timestamps):
+ out = frame_dir / f"scene_{base:08d}_f{i:02d}.jpg"
+ try:
+ extract_frame(video, ts, out, crop_subtitle=False, max_width=max_width)
+ paths.append(out)
+ except Exception as exc: # noqa: BLE001
+ print(f" [warn] scene frame {i} failed @ {ts:.2f}s: {exc}", flush=True)
+ return paths
+
+
+@dataclass
+class SceneVisionResult:
+ description: str = ""
+
+
+def describe_scene_vision(
+ client: LlamaServer,
+ scene: Scene,
+ images: list[Path],
+ speaker_registry: str,
+) -> SceneVisionResult:
+ if not images:
+ return SceneVisionResult()
+ spk_in_scene = sorted({c.speaker for c in scene.cues if c.speaker})
+ spk_line = ", ".join(spk_in_scene) if spk_in_scene else "unknown"
+ prompt = (
+ f"You are given {len(images)} frame(s) from the SAME scene, in chronological "
+ f"order (first image ≈ early, last ≈ late).\n"
+ f"Dialogue speakers tagged in this scene: {spk_line}.\n"
+ + (f"Known speakers:\n{speaker_registry}\n" if speaker_registry else "")
+ + "\nTasks:\n"
+ "1. Count visible people; note apparent gender and age (young adult / middle / "
+ "elderly) from appearance.\n"
+ "2. Infer relationship ONLY when visually clear: romantic_couple, parent_child, "
+ "colleagues, friends, strangers, unknown.\n"
+ "3. Do NOT assume parent–child just because one person looks older — drama "
+ "romances often pair young-looking actors.\n"
+ "4. Note setting (home, restaurant, office…) briefly.\n\n"
+ "Return ONLY JSON:\n"
+ '{"scene_description":"<1-2 sentences>",'
+ '"visible_characters":"",'
+ '"likely_relationship":"romantic_couple|parent_child|colleagues|friends|'
+ 'strangers|unknown",'
+ '"relationship_confidence":"high|medium|low"}'
+ )
+ try:
+ result = client.chat_json(
+ prompt,
+ images=images,
+ system=SCENE_VISION_SYSTEM,
+ max_tokens=512,
+ )
+ except Exception as exc: # noqa: BLE001
+ print(f" [warn] scene vision describe failed: {exc}", flush=True)
+ return SceneVisionResult()
+ if not isinstance(result, dict):
+ return SceneVisionResult()
+ parts = [
+ result.get("scene_description", ""),
+ result.get("visible_characters", ""),
+ ]
+ rel = (result.get("likely_relationship") or "").strip()
+ conf = (result.get("relationship_confidence") or "").strip()
+ if rel and rel != "unknown":
+ parts.append(f"Relationship: {rel} ({conf})")
+ description = " | ".join(p.strip() for p in parts if p and str(p).strip())
+ return SceneVisionResult(description=description)
+
+
+def _apply_scene_vision_to_cues(scene: Scene, description: str) -> None:
+ if not description:
+ return
+ for c in scene.cues:
+ c.visual_context = description
+
+
+def apply_speaker_report(cues: list[Cue], path: Path) -> int:
+ """Restore speaker/gender/age tags from a previous .report.json (skip diarization)."""
+ data = json.loads(path.read_text(encoding="utf-8"))
+ by_index = {int(item["index"]): item for item in data if "index" in item}
+ by_start = {item.get("start"): item for item in data if item.get("start")}
+ tagged = 0
+ for c in cues:
+ item = by_index.get(c.index) or by_start.get(c.start_ts)
+ if not item:
+ continue
+ spk = (item.get("speaker") or "").strip()
+ if spk:
+ c.speaker = spk
+ c.speaker_gender = (item.get("speaker_gender") or "").strip()
+ c.speaker_age_group = (item.get("speaker_age_group") or "").strip()
+ tagged += 1
+ return tagged
+
+
+def _history_block(previous: list[Cue], limit: int = 8) -> str:
+ if not previous:
+ return "(none)"
+ rows = []
+ for c in previous[-limit:]:
+ spk = _speaker_tag(c)
+ prefix = f"[{spk}] " if spk else ""
+ rows.append(f"- {prefix}src: {c.source!r} | translated: {c.translated!r}")
+ return "\n".join(rows)
+
+
+def build_translation_prompt(
+ scene: Scene,
+ previous: list[Cue],
+ source_lang: str,
+ target_lang: str,
+ max_line_chars: int = 42,
+ max_lines: int = 2,
+ speaker_registry: str = "",
+ name_registry: str = "",
+ relationship_block: str = "",
+ scene_vision: str = "",
+ n_scene_images: int = 1,
+) -> str:
+ src = "the source language" if source_lang == "auto" else source_lang
+ lines = []
+ for c in scene.cues:
+ budget = c.char_budget or subtitle_char_budget(c.duration)
+ spk = _speaker_tag(c)
+ spk_field = f" | speaker: {spk}" if spk else ""
+ # The scene-level description is shown once in the SCENE VISION block, so
+ # only attach a per-cue visual when it adds something different (avoids
+ # repeating the same paragraph on every cue line).
+ vis = (
+ c.visual_context
+ if c.visual_context and c.visual_context != scene_vision
+ else ""
+ )
+ vis_field = f" | visual: {vis!r}" if vis else ""
+ lines.append(
+ f"- index={c.index} | {c.start_ts} --> {c.end_ts} | "
+ f"duration={c.duration:.1f}s | max_chars={budget}{spk_field} | "
+ f"source: {c.source!r}{vis_field}"
+ )
+ line_rule = (
+ "- Each cue is ONE single line — never use \\n or line breaks in translated.\n"
+ if max_lines <= 1
+ else f"- Up to {max_lines} lines per cue (use \\n between lines if needed).\n"
+ )
+ speaker_block = ""
+ if speaker_registry:
+ speaker_block = (
+ "\nSPEAKER GUIDE (from voice diarization — who is speaking):\n"
+ + speaker_registry
+ + "\n- Each cue's 'speaker' field tells you WHICH character is talking. "
+ "Labels: gender nam=male, nữ=female, trẻ em=child, chưa rõ=unknown; age "
+ "thiếu niên=teen, thanh niên=young adult, trung niên=middle-aged, lớn "
+ "tuổi=elderly.\n"
+ "- GENDER is the reliable anchor: keep each speaker's gender consistent and "
+ "use it to pick the right gendered Vietnamese address — nam → anh/cậu/chú/"
+ "ông/lão...; nữ → cô/chị/dì/bà...; never flip a character's gender between "
+ "cues.\n"
+ "- AGE group is only a ROUGH HINT from voice and is OFTEN WRONG. Trust the "
+ "ADDRESS MAP (if provided), scene images and dialogue over the age label.\n"
+ "- 'chưa rõ' gender or a missing age means it is uncertain: infer everything "
+ "from dialogue, names and the scene image instead of guessing blindly.\n"
+ "- Keep ONE consistent persona per speaker label across the whole film.\n"
+ )
+ name_block = name_registry or ""
+ rel_block = relationship_block or ""
+ vision_intro = (
+ f"{n_scene_images} frame(s) from this scene are attached (chronological order).\n"
+ if n_scene_images > 1
+ else "An image of the current scene is attached for visual context.\n"
+ )
+ scene_vision_block = ""
+ if scene_vision:
+ scene_vision_block = (
+ f"\nSCENE VISION (from video frames — trust over wrong age labels):\n"
+ f"{scene_vision}\n"
+ )
+ if rel_block:
+ pronoun_block = (
+ "PRONOUNS / REGISTER:\n"
+ "- Follow the ADDRESS MAP above EXACTLY for every listed speaker pair: use "
+ "the given self-term and other-term and NEVER swap to a different pronoun "
+ "pair between cues (do not flip anh/em <-> mẹ/con <-> chị/em <-> cô/cháu, "
+ "etc.).\n"
+ f"- For pairs NOT in the map, infer the natural {target_lang} forms of "
+ "address from the dialogue, names, gender and scene, and keep them "
+ "consistent.\n"
+ "- Keep character names EXACTLY as in NAME GUIDE / earlier lines — do not "
+ "drift spellings between cues.\n\n"
+ )
+ else:
+ pronoun_block = (
+ "PRONOUNS / REGISTER:\n"
+ f"- Choose natural, consistent {target_lang} forms of address for each "
+ "speaker pair based on the dialogue, names, gender and scene. Once chosen, "
+ "NEVER switch the pronoun pair for the same two people across cues.\n"
+ "- Keep character names EXACTLY as in NAME GUIDE / earlier lines — do not "
+ "drift spellings between cues.\n\n"
+ )
+ return (
+ f"Translate the following subtitle cues from {src} into {target_lang}.\n"
+ + vision_intro
+ + scene_vision_block
+ + "\nTRANSLATION PRIORITY (natural & concise spoken Vietnamese):\n"
+ "- Render the MEANING the way a Vietnamese person would actually say it aloud — "
+ "not the literal words. Smooth, everyday spoken phrasing.\n"
+ "- Be concise: drop filler, redundant subjects/pronouns and repeated words. A "
+ "shorter, natural line beats a long faithful-but-stiff one.\n"
+ "- Avoid translationese and awkward calques (e.g. 'thông qua việc', 'một cách', "
+ "doubled 'tôi ... tôi ...'); use natural particles (à, nhé, thôi, đấy, mà) where they fit.\n"
+ "- Keep names, honorifics, emotional beats, jokes and plot-critical details.\n"
+ f"{line_rule}"
+ "- Keep within max_chars when you can; if a natural line is shorter, leave it short.\n"
+ + speaker_block
+ + rel_block
+ + name_block
+ + pronoun_block
+ + "Previously translated cues (for consistency):\n"
+ + _history_block(previous)
+ + "\n\nCues to translate now:\n"
+ + "\n".join(lines)
+ + "\n\nReturn ONLY a JSON array, one object per cue, in this exact shape:\n"
+ '[{"index": , "translated": "", '
+ '"honorific_notes": "", '
+ '"timing_notes": ""}]'
+ )
+
+
+def build_shorten_prompt(cue: Cue, target_lang: str, max_line_chars: int) -> str:
+ over = _translated_char_count(cue.translated) - cue.char_budget
+ return (
+ f"The subtitle below is somewhat long for its on-screen time ({over} chars over guide).\n"
+ f"Duration: {cue.duration:.1f}s | max_chars guide: {cue.char_budget} | "
+ f"source: {cue.source!r}\n"
+ f"Current translation ({_translated_char_count(cue.translated)} chars): "
+ f"{cue.translated!r}\n\n"
+ f"Rewrite in {target_lang} only if needed — keep the same meaning, tone, and "
+ "honorifics. Trim filler words only; do NOT drop key content. Prefer staying close "
+ "to the current wording over aggressive compression.\n"
+ "One single line only — do not use \\n.\n"
+ "Return ONLY JSON: "
+ '{"index": '
+ f"{cue.index}, "
+ '"translated": "", '
+ '"timing_notes": ""}'
+ )
+
+
+# --------------------------------------------------------------------------- #
+# Passes
+# --------------------------------------------------------------------------- #
+def correct_scene(
+ client: LlamaServer,
+ scene: Scene,
+ video: Path,
+ frame_dir: Path,
+ source_lang: str,
+) -> None:
+ images: list[Path] = []
+ for c in scene.cues:
+ out = frame_dir / f"cue_{c.index:05d}_sub.jpg"
+ try:
+ extract_frame(video, c.midpoint, out, crop_subtitle=True)
+ images.append(out)
+ except Exception as exc: # noqa: BLE001
+ print(f" [warn] frame extract failed cue #{c.index}: {exc}", flush=True)
+ # Fall back to an empty/missing frame: skip this image slot.
+ images.append(out if out.exists() else _placeholder_frame(frame_dir))
+
+ prompt = build_correction_prompt(scene, source_lang)
+ try:
+ result = client.chat_json(
+ prompt, images=images, system=CORRECTION_SYSTEM, max_tokens=1536
+ )
+ except Exception as exc: # noqa: BLE001
+ print(f" [warn] correction failed: {exc}", flush=True)
+ for c in scene.cues:
+ c.corrected_source = c.text
+ return
+
+ by_index = {}
+ if isinstance(result, list):
+ for item in result:
+ if isinstance(item, dict) and "cue_index" in item:
+ by_index[int(item["cue_index"])] = item
+ for c in scene.cues:
+ item = by_index.get(c.index, {})
+ c.corrected_source = (item.get("corrected_source") or c.text).strip()
+ c.was_corrected = bool(item.get("was_corrected", False))
+ c.ocr_text = (item.get("ocr_text") or "").strip()
+ c.visual_context = (item.get("visual_context") or "").strip()
+ c.correction_reason = (item.get("correction_reason") or "").strip()
+
+
+_PLACEHOLDER: Path | None = None
+
+
+def _placeholder_frame(frame_dir: Path) -> Path:
+ global _PLACEHOLDER
+ if _PLACEHOLDER and _PLACEHOLDER.exists():
+ return _PLACEHOLDER
+ out = frame_dir / "_placeholder.jpg"
+ _run_ffmpeg(
+ ["-f", "lavfi", "-i", "color=c=black:s=64x64", "-frames:v", "1", "-y", str(out)]
+ )
+ _PLACEHOLDER = out
+ return out
+
+
+def translate_scene(
+ client: LlamaServer,
+ scene: Scene,
+ previous: list[Cue],
+ video: Path,
+ frame_dir: Path,
+ source_lang: str,
+ target_lang: str,
+ use_scene_image: bool,
+ max_line_chars: int,
+ max_lines: int,
+ speaker_registry: str = "",
+ name_registry: NameRegistry | None = None,
+ address_registry: AddressRegistry | None = None,
+ scene_frames: int = 3,
+ use_scene_vision: bool = True,
+) -> None:
+ images: list[Path] = []
+ scene_vision = ""
+ if use_scene_image:
+ images = extract_scene_frames(
+ video, scene, frame_dir, n_frames=scene_frames
+ )
+ if use_scene_vision and images:
+ vision = describe_scene_vision(
+ client, scene, images, speaker_registry
+ )
+ scene_vision = vision.description
+ _apply_scene_vision_to_cues(scene, scene_vision)
+
+ prompt = build_translation_prompt(
+ scene,
+ previous,
+ source_lang,
+ target_lang,
+ max_line_chars=max_line_chars,
+ max_lines=max_lines,
+ speaker_registry=speaker_registry,
+ name_registry=name_registry.build_prompt_block() if name_registry else "",
+ relationship_block=(
+ address_registry.build_prompt_block(scene)
+ if address_registry
+ else ""
+ ),
+ scene_vision=scene_vision,
+ n_scene_images=len(images),
+ )
+ try:
+ result = client.chat_json(
+ prompt, images=images or None, system=TRANSLATION_SYSTEM, max_tokens=1536
+ )
+ except Exception as exc: # noqa: BLE001
+ print(f" [warn] translation failed: {exc}", flush=True)
+ for c in scene.cues:
+ c.translated = c.source
+ return
+
+ by_index = {}
+ if isinstance(result, list):
+ for item in result:
+ if isinstance(item, dict) and "index" in item:
+ by_index[int(item["index"])] = item
+ for c in scene.cues:
+ item = by_index.get(c.index, {})
+ c.translated = _normalize_subtitle_text(item.get("translated") or c.source)
+ c.honorific_notes = (item.get("honorific_notes") or "").strip()
+ c.timing_notes = (item.get("timing_notes") or "").strip()
+ if name_registry is not None:
+ name_registry.observe_cue(c)
+
+
+def shorten_overbudget_cues(
+ client: LlamaServer,
+ cues: list[Cue],
+ target_lang: str,
+ max_line_chars: int,
+ *,
+ slack_chars: int = 8,
+) -> int:
+ """Second pass: tighten only cues clearly over the char budget."""
+ fixed = 0
+ for c in cues:
+ if not c.translated or not c.char_budget:
+ continue
+ over = _translated_char_count(c.translated) - c.char_budget
+ if over <= slack_chars:
+ continue
+ prompt = build_shorten_prompt(c, target_lang, max_line_chars)
+ try:
+ result = client.chat_json(
+ prompt, system=TRANSLATION_SYSTEM, max_tokens=512
+ )
+ except Exception as exc: # noqa: BLE001
+ print(f" [warn] shorten failed cue #{c.index}: {exc}", flush=True)
+ continue
+ item = result if isinstance(result, dict) else {}
+ if isinstance(result, list) and result:
+ item = result[0] if isinstance(result[0], dict) else {}
+ shorter = (item.get("translated") or "").strip()
+ shorter = _normalize_subtitle_text(shorter)
+ if shorter and _translated_char_count(shorter) < _translated_char_count(c.translated):
+ c.translated = shorter
+ note = (item.get("timing_notes") or "").strip()
+ if note:
+ c.timing_notes = note
+ fixed += 1
+ return fixed
+
+
+# --------------------------------------------------------------------------- #
+# Pipeline
+# --------------------------------------------------------------------------- #
+def _plog(args: argparse.Namespace, msg: str) -> None:
+ log_fn = getattr(args, "log_fn", None)
+ if log_fn:
+ log_fn(msg)
+ else:
+ print(msg, flush=True)
+
+
+def _should_stop(args: argparse.Namespace) -> bool:
+ stop_check = getattr(args, "stop_check", None)
+ return bool(stop_check and stop_check())
+
+
+def run_pipeline(args: argparse.Namespace) -> None:
+ video = Path(args.video)
+ input_srt = Path(args.input_srt)
+ output_srt = Path(args.output_srt)
+ report_path = Path(args.report) if args.report else output_srt.with_suffix(
+ output_srt.suffix + ".report.json"
+ )
+
+ if not video.exists():
+ raise SystemExit(f"Video not found: {video}")
+ if not input_srt.exists():
+ raise SystemExit(f"SRT not found: {input_srt}")
+
+ server_bin = Path(args.llama_server)
+ model = Path(args.model)
+ mmproj = Path(args.mmproj)
+ if not args.server_url:
+ if not server_bin.exists():
+ raise SystemExit(f"llama-server not found: {server_bin}")
+ if not model.exists():
+ raise SystemExit(f"Model not found: {model}")
+ if not mmproj.exists():
+ raise SystemExit(f"mmproj not found: {mmproj}")
+
+ cues = parse_srt(input_srt)
+ if args.limit and args.limit > 0:
+ cues = cues[: args.limit]
+ if not cues:
+ raise SystemExit("No cues found in SRT.")
+ _plog(args, f"Loaded {len(cues)} cues from {input_srt.name}")
+
+ speaker_report = getattr(args, "speaker_report", "") or ""
+ if speaker_report:
+ report_src = Path(speaker_report)
+ if not report_src.exists():
+ raise SystemExit(f"Speaker report not found: {report_src}")
+ n_tagged = apply_speaker_report(cues, report_src)
+ _plog(args, f" Loaded speaker tags from report: {n_tagged}/{len(cues)} cue(s).")
+
+ speaker_genders: dict[str, str] = {}
+ if getattr(args, "diarize", False):
+ _plog(args, "Pass 0: speaker diarization (pyannote) — who speaks each cue...")
+ t0 = time.time()
+ try:
+ from diarize_audio import diarize_and_tag_cues
+
+ speaker_genders = diarize_and_tag_cues(
+ video,
+ cues,
+ hf_token=getattr(args, "hf_token", None),
+ num_speakers=getattr(args, "num_speakers", 0) or None,
+ min_speakers=getattr(args, "min_speakers", 0) or None,
+ max_speakers=getattr(args, "max_speakers", 0) or None,
+ device=getattr(args, "diarize_device", "auto"),
+ detect_gender=not getattr(args, "no_gender", False),
+ gender_method=getattr(args, "gender_method", "auto"),
+ merge_speakers=getattr(args, "merge_speakers", True),
+ merge_threshold=getattr(args, "merge_speaker_threshold", 0.70),
+ log=lambda m: _plog(args, m),
+ )
+ n_spk = len({c.speaker for c in cues if c.speaker})
+ _plog(
+ args,
+ f" Diarization done in {time.time() - t0:.0f}s "
+ f"({n_spk} speaker(s) tagged).",
+ )
+ except Exception as exc: # noqa: BLE001
+ _plog(args, f" [warn] diarization failed, continuing without it: {exc}")
+
+ scenes = group_scenes(
+ cues, args.scene_max_gap, args.scene_max_cues, args.scene_max_dur
+ )
+ _plog(args, f"Grouped into {len(scenes)} scene(s).")
+
+ _apply_char_budgets(
+ cues, args.chars_per_sec, args.max_line_chars, args.max_lines
+ )
+
+ draft = Path(args.model_draft) if args.model_draft else None
+ client = LlamaServer(
+ server_bin=server_bin,
+ model=model,
+ mmproj=mmproj,
+ ngl=args.ngl,
+ ctx_size=args.ctx,
+ port=args.port,
+ draft=draft,
+ use_mtp=not args.no_mtp,
+ temperature=args.temp,
+ existing_url=args.server_url,
+ )
+ client.start()
+
+ tmp = tempfile.mkdtemp(prefix="gemma_srt_")
+ frame_dir = Path(tmp)
+
+ try:
+ if not args.skip_correction:
+ _plog(args, "Pass 1/2: OCR + correcting source text (vision, per scene)...")
+ t0 = time.time()
+ for i, scene in enumerate(scenes, start=1):
+ if _should_stop(args):
+ _plog(args, "Stopped by user during correction pass.")
+ return
+ rng = f"#{scene.cues[0].index}-#{scene.cues[-1].index}"
+ _plog(args, f" [{i}/{len(scenes)}] scene {rng} @ {scene.start:.1f}s")
+ if getattr(args, "progress_fn", None):
+ args.progress_fn(i, len(scenes), "correction")
+ correct_scene(client, scene, video, frame_dir, args.source_lang)
+ n_fixed = sum(1 for c in cues if c.was_corrected)
+ _plog(args, f" Correction done in {time.time() - t0:.0f}s ({n_fixed} fixed).")
+ else:
+ for c in cues:
+ c.corrected_source = c.text
+
+ if args.merge_pairs:
+ before = len(cues)
+ cues, n_merge = merge_adjacent_pairs(
+ cues,
+ max_gap=args.merge_max_gap,
+ max_combined_duration=args.merge_max_duration,
+ max_combined_chars=args.merge_max_chars,
+ fragment_chars=args.merge_fragment_chars,
+ )
+ _plog(
+ args,
+ f"Merged {n_merge} fragment pair(s): {before} -> {len(cues)} cue(s) "
+ f"(same-speaker, sentence-aware).",
+ )
+ scenes = group_scenes(
+ cues, args.scene_max_gap, args.scene_max_cues, args.scene_max_dur
+ )
+ _apply_char_budgets(
+ cues, args.chars_per_sec, args.max_line_chars, args.max_lines
+ )
+
+ _plog(args, "Pass 2/2: translating with scene context + honorifics...")
+ speaker_registry = build_speaker_registry(cues)
+ if speaker_registry:
+ _plog(args, " Using speaker registry for consistent pronouns:")
+ for row in speaker_registry.splitlines():
+ _plog(args, f" {row}")
+ address_registry = AddressRegistry()
+ if getattr(args, "address_resolution", True):
+ _plog(
+ args,
+ " Resolving global forms of address (xưng hô) per speaker pair...",
+ )
+ t_addr = time.time()
+ address_registry = resolve_address_map(
+ client,
+ cues,
+ args.source_lang,
+ temperature=getattr(args, "address_temp", 0.2),
+ log=lambda m: _plog(args, m),
+ )
+ if address_registry.has_map:
+ _plog(
+ args,
+ f" Address map resolved in {time.time() - t_addr:.0f}s:",
+ )
+ for line in address_registry.describe_lines():
+ _plog(args, f" {line}")
+ else:
+ _plog(args, " No conversing speaker pairs resolved (continuing).")
+ name_registry = NameRegistry.from_cues(cues)
+ if name_registry.has_names:
+ _plog(
+ args,
+ f" Tracking {len(name_registry.source_tokens)} source name token(s) "
+ "for spelling consistency.",
+ )
+ scene_frames = max(1, getattr(args, "scene_frames", 3))
+ use_scene_vision = not getattr(args, "no_scene_vision", False)
+ if not args.no_scene_image:
+ _plog(
+ args,
+ f" Scene vision: {scene_frames} frame(s)/scene"
+ + (" + describe pass" if use_scene_vision else ""),
+ )
+ t0 = time.time()
+ translated: list[Cue] = []
+ for i, scene in enumerate(scenes, start=1):
+ if _should_stop(args):
+ _plog(args, "Stopped by user during translation pass.")
+ return
+ rng = f"#{scene.cues[0].index}-#{scene.cues[-1].index}"
+ _plog(args, f" [{i}/{len(scenes)}] scene {rng}")
+ if getattr(args, "progress_fn", None):
+ args.progress_fn(i, len(scenes), "translation")
+ translate_scene(
+ client,
+ scene,
+ translated,
+ video,
+ frame_dir,
+ args.source_lang,
+ args.target_lang,
+ use_scene_image=not args.no_scene_image,
+ max_line_chars=args.max_line_chars,
+ max_lines=args.max_lines,
+ speaker_registry=speaker_registry,
+ name_registry=name_registry,
+ address_registry=address_registry,
+ scene_frames=scene_frames,
+ use_scene_vision=use_scene_vision,
+ )
+ translated.extend(scene.cues)
+ name_registry.reconcile_all(cues)
+ if name_registry.source_to_vn:
+ _plog(args, " Locked name spellings:")
+ for src, vn in sorted(name_registry.source_to_vn.items()):
+ _plog(args, f" {src} → {vn}")
+ _plog(args, f" Translation done in {time.time() - t0:.0f}s.")
+
+ if args.shorten:
+ over = sum(
+ 1
+ for c in cues
+ if c.translated
+ and c.char_budget
+ and _translated_char_count(c.translated) > c.char_budget + args.shorten_slack
+ )
+ within = sum(
+ 1
+ for c in cues
+ if c.translated
+ and c.char_budget
+ and _translated_char_count(c.translated) <= c.char_budget + args.shorten_slack
+ )
+ if over:
+ _plog(
+ args,
+ f"Pass 2b: {within} cue(s) already within budget (skip); "
+ f"shortening {over} over-budget cue(s)...",
+ )
+ t0 = time.time()
+ n_fixed = shorten_overbudget_cues(
+ client,
+ cues,
+ args.target_lang,
+ args.max_line_chars,
+ slack_chars=args.shorten_slack,
+ )
+ still_over = sum(
+ 1
+ for c in cues
+ if c.translated
+ and c.char_budget
+ and _translated_char_count(c.translated) > c.char_budget
+ )
+ _plog(
+ args,
+ f" Shortened {n_fixed} cue(s) in {time.time() - t0:.0f}s "
+ f"({still_over} still over budget).",
+ )
+ else:
+ _plog(args, f"Pass 2b: all {within} cue(s) within budget — skip shorten.")
+
+ within = sum(
+ 1
+ for c in cues
+ if c.translated
+ and (
+ not c.char_budget
+ or _translated_char_count(c.translated) <= c.char_budget
+ )
+ )
+ _plog(args, f" Timing fit: {within}/{len(cues)} cues within char budget.")
+
+ write_srt(cues, output_srt, use_translation=True)
+ corrected_path = output_srt.with_suffix(".corrected" + output_srt.suffix)
+ write_srt(cues, corrected_path, use_translation=False)
+ write_report(cues, report_path)
+
+ _plog(args, "Done.")
+ _plog(args, f" Translated SRT: {output_srt}")
+ _plog(args, f" Corrected source: {corrected_path}")
+ _plog(args, f" Report JSON: {report_path}")
+ finally:
+ if client.owns_server:
+ client.stop()
+
+
+# --------------------------------------------------------------------------- #
+# CLI
+# --------------------------------------------------------------------------- #
+def build_parser() -> argparse.ArgumentParser:
+ p = argparse.ArgumentParser(
+ description="Correct and translate SRT using Gemma 4 vision + scene context "
+ "(persistent llama-server)."
+ )
+ p.add_argument("--video", required=True, help="Input video file")
+ p.add_argument("--input-srt", required=True, help="Source SRT file")
+ p.add_argument("--output-srt", required=True, help="Output translated SRT path")
+ p.add_argument("--report", default="", help="JSON report path (default: