Download chat_template.py from FannyFa/FannyFa-Model-V1: direct link, hf CLI and curl.
- Browser
- Download file 3.83 kB
-
https://huggingface.co/FannyFa/FannyFa-Model-V1/resolve/main/chat_template.py
- Command line
-
hf download hf://FannyFa/FannyFa-Model-V1/chat_template.py
-
curl -L -o chat_template.py https://huggingface.co/FannyFa/FannyFa-Model-V1/resolve/main/chat_template.py
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 |