Files
live-app-translator/scripts/finetune_mt.py
EJClaw 5596f905d1 fix: Seed-X 프롬프트 형식과 티어별 용어집 전략 수정
4티어(Seed-X-PPO-7B) 번역 경로가 모델 카드 요구사항을 어기고 있었다.

- 프롬프트 끝의 `<ko>` 등 대상 언어 태그가 빠져 있었다. PPO 학습에 쓰인
  신호라 없으면 번역 품질이 흔들린다.
- Seed-X 는 chat template 없는 번역 전용 completion 모델인데 "구어체로
  자연스럽게" 같은 지시문 래퍼를 씌우고 있었다. 학습 분포를 벗어난다.
- 그 결과 MTSpec.supports_prompt_glossary=True 가 사실과 달랐다.
  Seed-X 는 용어집 지시문을 이해하지 못하므로 플레이스홀더 치환을 써야 한다.

수정
- PromptStyle(NONE/SEEDX/INSTRUCT) 도입, supports_prompt_glossary 를
  prompt_style 에서 파생시켜 둘이 어긋날 수 없게 함
- build_seedx_prompt() 로 모델 카드 형식을 분리 (지시문 주입 불가)
- INSTRUCT 경로는 chat template 사용, Qwen3 thinking 모드는 끔
- LANGUAGES 에 seedx 태그 명시
- finetune_mt.py 가 티어의 prompt_style 을 따라가게 해 학습/추론 프롬프트 일치
- AWQ Int4 로드 실패 시 autoawq 설치 안내를 담은 오류 메시지

검증: pytest 52개 통과 (프롬프트 회귀 테스트 13개 추가), ruff clean

근거: https://huggingface.co/ByteDance-Seed/Seed-X-PPO-7B

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-09-21 10:56:49 +09:00

270 lines
9.5 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 ( # noqa: E402
MTBackend,
PromptStyle,
get_tier,
ordered_tiers,
)
from hearo.models.translator import build_seedx_prompt # 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, style: PromptStyle = PromptStyle.INSTRUCT) -> tuple[str, str]:
"""LLM 학습용 (프롬프트, 정답) 쌍.
**추론 때 쓰는 프롬프트와 반드시 같아야 한다.** 형식이 어긋나면 학습은
정상적으로 끝나지만 실제 번역에서는 효과가 거의 나오지 않는다. 그래서
Seed-X 형식은 translator.py 의 함수를 그대로 재사용한다.
"""
if style is PromptStyle.SEEDX:
prompt = build_seedx_prompt(
item["source"], item["source_lang"], item["target_lang"]
)
return prompt, "\n" + item["target"]
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,
style: PromptStyle = PromptStyle.INSTRUCT):
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, style)
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,
tier.mt.prompt_style,
)
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())