FannyFa-Model-V1 / chat_template.py
FannyFa's picture
Upload 16 files
3eecd6b verified
Raw History Blame Contribute Delete
3.83 kB
"""
Chat Template Helper
Auto-detect format: native chat template (Qwen/Llama/Mistral) atau simple (GPT-2).
Dipakai oleh chat.py, rag_module.py, dan trainer.py.
"""
FORMAT_NATIVE = "native"
FORMAT_SIMPLE = "simple"
DEFAULT_SYSTEM_PROMPT = "Anda adalah asisten AI yang ramah dan informatif."
def detect_format(tokenizer) -> str:
"""Deteksi format yang didukung tokenizer."""
if getattr(tokenizer, "chat_template", None):
return FORMAT_NATIVE
return FORMAT_SIMPLE
def format_inference_prompt(tokenizer, user_message, system_message=None, history=None):
"""Format prompt untuk inference/chat."""
fmt = detect_format(tokenizer)
system_message = system_message or DEFAULT_SYSTEM_PROMPT
if fmt == FORMAT_NATIVE:
messages = [{"role": "system", "content": system_message}]
if history:
for turn in history[-5:]:
u = (turn.get("user") or "").strip()
a = (turn.get("ai") or "").strip()
if u and a:
messages.append({"role": "user", "content": u})
messages.append({"role": "assistant", "content": a})
messages.append({"role": "user", "content": user_message})
return tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=True
)
# Fallback: format lama "User: / AI:"
parts = [system_message]
if history:
for turn in history[-2:]:
u = (turn.get("user") or "").strip()
a = (turn.get("ai") or "").strip()
if u and a:
parts.append(f"User: {u}")
parts.append(f"AI: {a}")
parts.append(f"User: {user_message}")
parts.append("AI:")
return "\n".join(parts)
def format_training_sample(tokenizer, user_message, ai_response, system_message=None):
"""Format 1 Q&A untuk training."""
fmt = detect_format(tokenizer)
system_message = system_message or DEFAULT_SYSTEM_PROMPT
if fmt == FORMAT_NATIVE:
messages = [
{"role": "system", "content": system_message},
{"role": "user", "content": user_message},
{"role": "assistant", "content": ai_response},
]
return tokenizer.apply_chat_template(
messages, tokenize=False, add_generation_prompt=False
)
eos = tokenizer.eos_token or ""
return f"User: {user_message}\nAI: {ai_response}{eos}"
def tokenize_for_training(tokenizer, text, max_length):
"""Tokenize + prompt masking untuk training."""
enc = tokenizer(
text,
truncation=True,
max_length=max_length,
padding=False,
return_tensors=None,
)
ids = enc["input_ids"]
labels = list(ids)
fmt = detect_format(tokenizer)
marker = "<|im_start|>assistant\n" if fmt == FORMAT_NATIVE else "AI:"
marker_ids = tokenizer.encode(marker, add_special_tokens=False)
mlen = len(marker_ids)
if mlen > 0 and len(ids) > mlen:
pos = -1
for i in range(len(ids) - mlen, -1, -1):
if ids[i:i + mlen] == marker_ids:
pos = i + mlen
break
if pos > 0:
for i in range(min(pos, len(labels))):
labels[i] = -100
return {
"input_ids": ids,
"attention_mask": enc["attention_mask"],
"labels": labels,
}
def tokenize_batch_for_training(examples, tokenizer, max_length):
"""Tokenize batch (untuk dataset.map)."""
result = {"input_ids": [], "attention_mask": [], "labels": []}
for text in examples["text"]:
out = tokenize_for_training(tokenizer, text, max_length)
result["input_ids"].append(out["input_ids"])
result["attention_mask"].append(out["attention_mask"])
result["labels"].append(out["labels"])
return result