fix: apply chat template in DPO tokenization
- dpo_tokenize now uses tokenizer.apply_chat_template to match SFT format - Prompt rendered with add_generation_prompt=True - Chosen/rejected appended as assistant turn - Remove leftover dead code from _extract_text - Update tests to mock apply_chat_template
This commit is contained in:
+60
-23
@@ -39,48 +39,85 @@ def dpo_tokenize(
|
||||
record: dict,
|
||||
tokenizer,
|
||||
max_len: int = 2048,
|
||||
pad_id: int = 2,
|
||||
) -> Optional[dict]:
|
||||
"""Tokenize one DPO record into chosen/rejected + masks.
|
||||
|
||||
Pure processor function (HF ``datasets.map`` style):
|
||||
``record -> dict_of_lists``. Each value is a flat list of ints/bools.
|
||||
Applies the tokenizer's chat template so token sequences match the
|
||||
SFT checkpoint's format. Prompt is rendered with
|
||||
``add_generation_prompt=True``; chosen/rejected are appended as a
|
||||
single assistant turn.
|
||||
|
||||
No packing, no ``position_ids`` — DPO sequences are independent and
|
||||
the model defaults to ``arange(0, seq_len)``.
|
||||
Accepts:
|
||||
|
||||
- Flat: ``{"prompt": str, "chosen": str, "rejected": str}``
|
||||
- Conv: ``{"prompt": [{role, content}, ...], "chosen": [...], ...}``
|
||||
- Legacy: ``{"input": str, "chosen": str, "rejected": str}``
|
||||
|
||||
No packing, no ``position_ids`` — DPO sequences are independent.
|
||||
"""
|
||||
inp = record.get("input")
|
||||
chosen_text = record.get("chosen")
|
||||
rejected_text = record.get("rejected")
|
||||
if inp is None or chosen_text is None or rejected_text is None:
|
||||
prompt = record.get("prompt") or record.get("input")
|
||||
chosen = record.get("chosen")
|
||||
rejected = record.get("rejected")
|
||||
if prompt is None or chosen is None or rejected is None:
|
||||
return None
|
||||
|
||||
in_ids = tokenizer.encode(inp, add_special_tokens=True)
|
||||
ch_ids = tokenizer.encode(chosen_text, add_special_tokens=False)
|
||||
re_ids = tokenizer.encode(rejected_text, add_special_tokens=False)
|
||||
prompt_messages = _to_messages(prompt)
|
||||
chosen_text = _extract_text(chosen)
|
||||
rejected_text = _extract_text(rejected)
|
||||
if chosen_text is None or rejected_text is None:
|
||||
return None
|
||||
chosen_messages = prompt_messages + [{"role": "assistant", "content": chosen_text}]
|
||||
rejected_messages = prompt_messages + [
|
||||
{"role": "assistant", "content": rejected_text}
|
||||
]
|
||||
|
||||
full_ch = (in_ids + ch_ids)[:max_len]
|
||||
full_re = (in_ids + re_ids)[:max_len]
|
||||
prompt_ids = tokenizer.apply_chat_template(
|
||||
prompt_messages, tokenize=True, add_generation_prompt=True
|
||||
)
|
||||
ch_ids = tokenizer.apply_chat_template(
|
||||
chosen_messages, tokenize=True, add_generation_prompt=False
|
||||
)
|
||||
re_ids = tokenizer.apply_chat_template(
|
||||
rejected_messages, tokenize=True, add_generation_prompt=False
|
||||
)
|
||||
|
||||
max_record_len = max(len(full_ch), len(full_re))
|
||||
ch_pad = full_ch + [pad_id] * (max_record_len - len(full_ch))
|
||||
re_pad = full_re + [pad_id] * (max_record_len - len(full_re))
|
||||
full_ch = ch_ids[:max_len]
|
||||
full_re = re_ids[:max_len]
|
||||
|
||||
ch_mask = [0] * len(in_ids) + [1] * len(ch_ids)
|
||||
prompt_len = min(len(prompt_ids), max_len)
|
||||
ch_mask = [0] * prompt_len + [1] * max(0, len(full_ch) - prompt_len)
|
||||
ch_mask = ch_mask[:max_len]
|
||||
ch_mask += [0] * (max_record_len - len(ch_mask))
|
||||
re_mask = [0] * len(in_ids) + [1] * len(re_ids)
|
||||
re_mask = [0] * prompt_len + [1] * max(0, len(full_re) - prompt_len)
|
||||
re_mask = re_mask[:max_len]
|
||||
re_mask += [0] * (max_record_len - len(re_mask))
|
||||
|
||||
return {
|
||||
"chosen": ch_pad,
|
||||
"rejected": re_pad,
|
||||
"chosen": full_ch,
|
||||
"rejected": full_re,
|
||||
"chosen_mask": ch_mask,
|
||||
"rejected_mask": re_mask,
|
||||
}
|
||||
|
||||
|
||||
def _to_messages(value) -> list:
|
||||
"""Accept str or conversation list; return message list."""
|
||||
if isinstance(value, str):
|
||||
return [{"role": "user", "content": value}]
|
||||
if isinstance(value, list):
|
||||
return value
|
||||
return [{"role": "user", "content": str(value)}]
|
||||
|
||||
|
||||
def _extract_text(value) -> Optional[str]:
|
||||
"""Accept str or conversation list; return plain text."""
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, str):
|
||||
return value
|
||||
if isinstance(value, list):
|
||||
return "".join(m.get("content", "") for m in value if isinstance(m, dict))
|
||||
return None
|
||||
|
||||
|
||||
def dpo_processor(
|
||||
record: dict,
|
||||
tokenizer,
|
||||
|
||||
@@ -917,16 +917,23 @@ def test_dpo_tokenize_pure_function():
|
||||
"""dpo_tokenize returns flat lists with correct mask alignment."""
|
||||
|
||||
class FakeTokenizer:
|
||||
def encode(self, text, add_special_tokens=True):
|
||||
return [len(text)] if add_special_tokens else [len(text) + 1]
|
||||
def apply_chat_template(
|
||||
self, messages, tokenize=True, add_generation_prompt=True
|
||||
):
|
||||
ids = []
|
||||
for m in messages:
|
||||
ids.append(len(m["content"]))
|
||||
ids.append(-1)
|
||||
if add_generation_prompt:
|
||||
ids.append(99)
|
||||
return ids
|
||||
|
||||
record = {"input": "ab", "chosen": "xyz", "rejected": "w"}
|
||||
result = dpo_tokenize(record, FakeTokenizer(), max_len=64, pad_id=0)
|
||||
record = {"prompt": "ab", "chosen": "xyz", "rejected": "w"}
|
||||
result = dpo_tokenize(record, FakeTokenizer(), max_len=64)
|
||||
|
||||
assert set(result.keys()) == {"chosen", "rejected", "chosen_mask", "rejected_mask"}
|
||||
assert len(result["chosen"]) == len(result["chosen_mask"])
|
||||
assert len(result["rejected"]) == len(result["rejected_mask"])
|
||||
assert len(result["chosen"]) == len(result["rejected"])
|
||||
|
||||
assert result["chosen_mask"][0] == 0
|
||||
assert any(m == 1 for m in result["chosen_mask"])
|
||||
@@ -937,12 +944,14 @@ def test_dpo_tokenize_malformed_record():
|
||||
"""dpo_tokenize returns None for missing fields."""
|
||||
|
||||
class FakeTokenizer:
|
||||
def encode(self, text, add_special_tokens=True):
|
||||
def apply_chat_template(
|
||||
self, messages, tokenize=True, add_generation_prompt=True
|
||||
):
|
||||
return [1]
|
||||
|
||||
assert dpo_tokenize({}, FakeTokenizer()) is None
|
||||
assert dpo_tokenize({"input": "a"}, FakeTokenizer()) is None
|
||||
assert dpo_tokenize({"input": "a", "chosen": "b"}, FakeTokenizer()) is None
|
||||
assert dpo_tokenize({"prompt": "a"}, FakeTokenizer()) is None
|
||||
assert dpo_tokenize({"prompt": "a", "chosen": "b"}, FakeTokenizer()) is None
|
||||
|
||||
|
||||
def test_dpo_jsonl_lazy_load(base_test_env):
|
||||
|
||||
Reference in New Issue
Block a user