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>
This commit is contained in:
@@ -21,7 +21,13 @@ 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
|
||||
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"]
|
||||
@@ -52,8 +58,19 @@ def load_pairs(path: Path) -> list[dict]:
|
||||
return pairs
|
||||
|
||||
|
||||
def build_prompt(item: dict) -> tuple[str, str]:
|
||||
"""LLM 학습용 (프롬프트, 정답) 쌍. translator.py 의 추론 프롬프트와 맞춰야 한다."""
|
||||
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 = (
|
||||
@@ -65,7 +82,8 @@ def build_prompt(item: dict) -> tuple[str, str]:
|
||||
return prompt, " " + item["target"]
|
||||
|
||||
|
||||
def train_llm(pairs, repo, output: Path, epochs: int, lr: float, rank: int, batch: int):
|
||||
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
|
||||
@@ -82,7 +100,7 @@ def train_llm(pairs, repo, output: Path, epochs: int, lr: float, rank: int, batc
|
||||
tokenizer.pad_token = tokenizer.eos_token
|
||||
|
||||
def encode(item):
|
||||
prompt, answer = build_prompt(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)
|
||||
@@ -227,7 +245,10 @@ def main() -> int:
|
||||
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)
|
||||
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)
|
||||
|
||||
Reference in New Issue
Block a user