프로그램별 오디오를 캡처해 로컬 GPU에서 음성인식→번역하고 화면 위 자막으로 보여주는 데스크톱 앱. 한/영/일/중 4개 언어. 구성 - audio: WASAPI 프로그램별 캡처(C++ 보조 프로그램) + 장치 루프백 폴백, 적응형 VAD 발화 분할 - models: 속도~품질 5단계 티어, faster-whisper + CTranslate2/LLM 2백엔드, 용어집(플레이스홀더 보호 + 프롬프트 주입) - core: Qt 비의존 파이프라인 엔진 (캡처/분할/추론 3스레드, 큐 연결) - ui: 사이드바 5화면 + 무테두리 항상위 자막 오버레이, 자체 다크 테마 모델 선정 근거는 docs/MODELS.md, 추가학습 가능 여부와 방법은 docs/FINETUNING.md 참고. 검증: pytest 39개 통과 (GPU·오디오 장치 없이 실행), ruff clean Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
249 lines
8.9 KiB
Python
249 lines
8.9 KiB
Python
#!/usr/bin/env python3
|
|
"""번역 모델 추가학습 (게임·방송 용어).
|
|
|
|
사용법:
|
|
python scripts/finetune_mt.py --data data/game.jsonl --tier ultimate \
|
|
--output ./adapters/game-ko --epochs 3
|
|
|
|
데이터 형식 (JSONL, 한 줄에 한 쌍):
|
|
{"source": "...", "target": "...", "source_lang": "en", "target_lang": "ko"}
|
|
|
|
자세한 배경은 docs/FINETUNING.md 참고.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
import sys
|
|
from pathlib import Path
|
|
|
|
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
|
|
|
|
from hearo.constants import LANGUAGES, models_dir # noqa: E402
|
|
from hearo.models.tiers import MTBackend, get_tier, ordered_tiers # noqa: E402
|
|
|
|
#: NLLB(seq2seq)와 Qwen(causal LM)에서 공통으로 존재하는 attention 투영 이름
|
|
LORA_TARGETS = ["q_proj", "v_proj"]
|
|
|
|
|
|
def load_pairs(path: Path) -> list[dict]:
|
|
pairs = []
|
|
with path.open(encoding="utf-8") as fh:
|
|
for lineno, raw in enumerate(fh, 1):
|
|
raw = raw.strip()
|
|
if not raw:
|
|
continue
|
|
try:
|
|
item = json.loads(raw)
|
|
except json.JSONDecodeError as exc:
|
|
raise SystemExit(f"{path}:{lineno} JSON 오류: {exc}") from exc
|
|
for key in ("source", "target"):
|
|
if not item.get(key):
|
|
raise SystemExit(f"{path}:{lineno} '{key}' 가 비어 있습니다")
|
|
item.setdefault("source_lang", "en")
|
|
item.setdefault("target_lang", "ko")
|
|
for key in ("source_lang", "target_lang"):
|
|
if item[key] not in LANGUAGES:
|
|
raise SystemExit(f"{path}:{lineno} 지원하지 않는 언어: {item[key]}")
|
|
pairs.append(item)
|
|
if not pairs:
|
|
raise SystemExit(f"{path} 에 학습 데이터가 없습니다")
|
|
return pairs
|
|
|
|
|
|
def build_prompt(item: dict) -> tuple[str, str]:
|
|
"""LLM 학습용 (프롬프트, 정답) 쌍. translator.py 의 추론 프롬프트와 맞춰야 한다."""
|
|
src = LANGUAGES[item["source_lang"]]["english"]
|
|
tgt = LANGUAGES[item["target_lang"]]["english"]
|
|
prompt = (
|
|
f"Translate the following {src} text into {tgt}. "
|
|
"It is a live spoken line from a game or broadcast, so keep the tone "
|
|
"casual and natural. Output only the translation.\n\n"
|
|
f"{src}: {item['source']}\n{tgt}:"
|
|
)
|
|
return prompt, " " + item["target"]
|
|
|
|
|
|
def train_llm(pairs, repo, output: Path, epochs: int, lr: float, rank: int, batch: int):
|
|
import torch
|
|
from datasets import Dataset
|
|
from peft import LoraConfig, get_peft_model
|
|
from transformers import (
|
|
AutoModelForCausalLM,
|
|
AutoTokenizer,
|
|
DataCollatorForLanguageModeling,
|
|
Trainer,
|
|
TrainingArguments,
|
|
)
|
|
|
|
tokenizer = AutoTokenizer.from_pretrained(repo, cache_dir=str(models_dir()))
|
|
if tokenizer.pad_token is None:
|
|
tokenizer.pad_token = tokenizer.eos_token
|
|
|
|
def encode(item):
|
|
prompt, answer = build_prompt(item)
|
|
prompt_ids = tokenizer(prompt, add_special_tokens=False)["input_ids"]
|
|
answer_ids = tokenizer(answer, add_special_tokens=False)["input_ids"]
|
|
answer_ids.append(tokenizer.eos_token_id)
|
|
ids = (prompt_ids + answer_ids)[:512]
|
|
# 프롬프트 부분은 손실에서 제외해야 '번역 결과'만 학습된다.
|
|
labels = ([-100] * len(prompt_ids) + answer_ids)[:512]
|
|
return {"input_ids": ids, "labels": labels, "attention_mask": [1] * len(ids)}
|
|
|
|
dataset = Dataset.from_list([encode(p) for p in pairs])
|
|
|
|
model = AutoModelForCausalLM.from_pretrained(
|
|
repo,
|
|
cache_dir=str(models_dir()),
|
|
torch_dtype=torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16,
|
|
device_map="auto",
|
|
)
|
|
model.enable_input_require_grads()
|
|
model = get_peft_model(
|
|
model,
|
|
LoraConfig(
|
|
r=rank,
|
|
lora_alpha=rank * 2,
|
|
lora_dropout=0.05,
|
|
bias="none",
|
|
task_type="CAUSAL_LM",
|
|
target_modules=LORA_TARGETS,
|
|
),
|
|
)
|
|
model.print_trainable_parameters()
|
|
|
|
Trainer(
|
|
model=model,
|
|
args=TrainingArguments(
|
|
output_dir=str(output / "checkpoints"),
|
|
num_train_epochs=epochs,
|
|
per_device_train_batch_size=batch,
|
|
gradient_accumulation_steps=max(1, 8 // batch),
|
|
learning_rate=lr,
|
|
warmup_ratio=0.05,
|
|
logging_steps=20,
|
|
save_strategy="no",
|
|
gradient_checkpointing=True,
|
|
report_to=[],
|
|
),
|
|
train_dataset=dataset,
|
|
data_collator=DataCollatorForLanguageModeling(tokenizer, mlm=False),
|
|
).train()
|
|
|
|
model.save_pretrained(output)
|
|
tokenizer.save_pretrained(output)
|
|
|
|
|
|
def train_seq2seq(pairs, repo, output: Path, epochs: int, lr: float, rank: int, batch: int):
|
|
"""NLLB 계열. CTranslate2 변환 모델이 아니라 원본 HF 모델로 학습해야 한다."""
|
|
from datasets import Dataset
|
|
from peft import LoraConfig, get_peft_model
|
|
from transformers import (
|
|
AutoModelForSeq2SeqLM,
|
|
AutoTokenizer,
|
|
DataCollatorForSeq2Seq,
|
|
Seq2SeqTrainer,
|
|
Seq2SeqTrainingArguments,
|
|
)
|
|
|
|
tokenizer = AutoTokenizer.from_pretrained(repo, cache_dir=str(models_dir()))
|
|
|
|
def encode(item):
|
|
tokenizer.src_lang = LANGUAGES[item["source_lang"]]["nllb"]
|
|
tokenizer.tgt_lang = LANGUAGES[item["target_lang"]]["nllb"]
|
|
batch_enc = tokenizer(
|
|
item["source"], text_target=item["target"], truncation=True, max_length=256
|
|
)
|
|
return dict(batch_enc)
|
|
|
|
dataset = Dataset.from_list([encode(p) for p in pairs])
|
|
model = AutoModelForSeq2SeqLM.from_pretrained(repo, cache_dir=str(models_dir()))
|
|
model = get_peft_model(
|
|
model,
|
|
LoraConfig(
|
|
r=rank,
|
|
lora_alpha=rank * 2,
|
|
lora_dropout=0.05,
|
|
bias="none",
|
|
task_type="SEQ_2_SEQ_LM",
|
|
target_modules=LORA_TARGETS,
|
|
),
|
|
)
|
|
model.print_trainable_parameters()
|
|
|
|
Seq2SeqTrainer(
|
|
model=model,
|
|
args=Seq2SeqTrainingArguments(
|
|
output_dir=str(output / "checkpoints"),
|
|
num_train_epochs=epochs,
|
|
per_device_train_batch_size=batch,
|
|
learning_rate=lr,
|
|
warmup_ratio=0.05,
|
|
logging_steps=20,
|
|
save_strategy="no",
|
|
report_to=[],
|
|
),
|
|
train_dataset=dataset,
|
|
data_collator=DataCollatorForSeq2Seq(tokenizer, model=model),
|
|
).train()
|
|
|
|
model.save_pretrained(output)
|
|
tokenizer.save_pretrained(output)
|
|
|
|
|
|
def main() -> int:
|
|
parser = argparse.ArgumentParser(description="번역 모델 LoRA 추가학습")
|
|
parser.add_argument("--data", required=True, type=Path, help="JSONL 학습 데이터")
|
|
parser.add_argument(
|
|
"--tier",
|
|
default="ultimate",
|
|
choices=[t.key for t in ordered_tiers()],
|
|
help="어느 티어의 번역 모델을 학습할지 (기본: ultimate)",
|
|
)
|
|
parser.add_argument("--output", required=True, type=Path, help="어댑터 저장 폴더")
|
|
parser.add_argument("--epochs", type=int, default=3)
|
|
parser.add_argument("--lr", type=float, default=2e-4)
|
|
parser.add_argument("--rank", type=int, default=16)
|
|
parser.add_argument("--batch", type=int, default=2)
|
|
parser.add_argument(
|
|
"--base-model",
|
|
default="",
|
|
help="티어 기본값 대신 쓸 HF 모델 (NLLB는 CT2 변환본이 아닌 원본을 지정해야 함)",
|
|
)
|
|
args = parser.parse_args()
|
|
|
|
tier = get_tier(args.tier)
|
|
pairs = load_pairs(args.data)
|
|
args.output.mkdir(parents=True, exist_ok=True)
|
|
|
|
if tier.mt.finetune_ease >= 3:
|
|
print(
|
|
f"경고: '{tier.name}' 티어의 번역 모델은 추가학습에 적합하지 않습니다.\n"
|
|
" docs/FINETUNING.md 참고 — 'ultimate' 티어를 권장합니다.\n",
|
|
file=sys.stderr,
|
|
)
|
|
|
|
if tier.mt.backend is MTBackend.TRANSFORMERS:
|
|
repo = args.base_model or tier.mt.repo
|
|
print(f"LLM LoRA 학습: {repo} · {len(pairs)}쌍 · {args.epochs} epoch")
|
|
train_llm(pairs, repo, args.output, args.epochs, args.lr, args.rank, args.batch)
|
|
else:
|
|
# CT2 변환본에는 학습에 필요한 가중치가 없으므로 원본 리포로 바꾼다.
|
|
repo = args.base_model or _hf_original(tier.mt.repo)
|
|
print(f"seq2seq LoRA 학습: {repo} · {len(pairs)}쌍 · {args.epochs} epoch")
|
|
train_seq2seq(pairs, repo, args.output, args.epochs, args.lr, args.rank, args.batch)
|
|
|
|
print(f"\n완료. '모델' 화면의 추가학습 어댑터 칸에 입력하세요:\n {args.output.resolve()}")
|
|
return 0
|
|
|
|
|
|
def _hf_original(ct2_repo: str) -> str:
|
|
"""CTranslate2 변환 리포 이름에서 원본 facebook/ 리포를 추정한다."""
|
|
name = ct2_repo.split("/")[-1].replace("-ctranslate2", "")
|
|
return f"facebook/{name}"
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|