initial upload
This commit is contained in:
348
utils/trainer_dataloaders.py
Normal file
348
utils/trainer_dataloaders.py
Normal file
@@ -0,0 +1,348 @@
|
||||
# utils/trainer_dataloaders.py
|
||||
|
||||
import logging
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from datasets import load_dataset, Dataset
|
||||
from utils.dataset_helpers import load_ftpo_multi_dataset
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
def load_and_prepare_dataset(config: dict, experiment_run_dir: Path, tokenizer: "AutoTokenizer") -> Dataset | None:
|
||||
"""
|
||||
Loads and prepares the dataset based on the finetuning mode specified in the config.
|
||||
|
||||
Args:
|
||||
config (dict): The experiment configuration dictionary.
|
||||
experiment_run_dir (Path): The directory for the current experiment run.
|
||||
tokenizer (AutoTokenizer): The tokenizer to use for processing.
|
||||
|
||||
Returns:
|
||||
Dataset or None: The prepared Hugging Face dataset, or None if loading fails.
|
||||
"""
|
||||
mode = config.get("finetune_mode", "ftpo").lower()
|
||||
max_seq_length = config['finetune_max_seq_length']
|
||||
dpo_dataset_hf = None
|
||||
|
||||
if mode == "dpo":
|
||||
# full-sequence preference pairs: rejected is baseline; chosen is the generation made with antislop
|
||||
dataset_path = experiment_run_dir / "dpo_pairs_dataset.jsonl"
|
||||
if not dataset_path.is_file():
|
||||
logger.error(f"DPO dataset not found at {dataset_path}")
|
||||
return None
|
||||
|
||||
dpo_dataset_hf = load_dataset(
|
||||
"json",
|
||||
data_files=str(dataset_path),
|
||||
split="train"
|
||||
)
|
||||
|
||||
# ----------------------------------------------------------
|
||||
# discard rows whose prompt+continuation would overflow
|
||||
# ----------------------------------------------------------
|
||||
def _within_len(example):
|
||||
prompt_ids = tokenizer(example["prompt"],
|
||||
add_special_tokens=False).input_ids
|
||||
chosen_ids = tokenizer(example["chosen"],
|
||||
add_special_tokens=False).input_ids
|
||||
rejected_ids = tokenizer(example["rejected"],
|
||||
add_special_tokens=False).input_ids
|
||||
max_len = config['finetune_max_seq_length']
|
||||
return (
|
||||
len(prompt_ids) + len(chosen_ids) <= max_len
|
||||
and
|
||||
len(prompt_ids) + len(rejected_ids) <= max_len
|
||||
)
|
||||
|
||||
before = len(dpo_dataset_hf)
|
||||
dpo_dataset_hf = dpo_dataset_hf.filter(_within_len)
|
||||
after = len(dpo_dataset_hf)
|
||||
logger.info(f"DPO length filter: kept {after}/{before} examples "
|
||||
f"(max_seq_len = {config['finetune_max_seq_length']})")
|
||||
|
||||
if after == 0:
|
||||
raise ValueError("every DPO sample exceeded finetune_max_seq_length")
|
||||
|
||||
|
||||
dpo_dataset_hf = dpo_dataset_hf.shuffle(seed=config.get("finetune_shuffle_seed", 3407))
|
||||
max_train = config.get("finetune_max_train_examples")
|
||||
if isinstance(max_train, int) and max_train > 0 and len(dpo_dataset_hf) > max_train:
|
||||
dpo_dataset_hf = dpo_dataset_hf.select(range(max_train))
|
||||
logger.info(f"Capped training dataset to {max_train} examples.")
|
||||
|
||||
# ── filter malformed rows (prompt / chosen / rejected missing) ──
|
||||
req_cols = {"prompt", "chosen", "rejected"}
|
||||
before_len = len(dpo_dataset_hf)
|
||||
dpo_dataset_hf = dpo_dataset_hf.filter(
|
||||
lambda x: all(col in x and x[col] for col in req_cols)
|
||||
)
|
||||
after_len = len(dpo_dataset_hf)
|
||||
if after_len == 0:
|
||||
logger.error("All rows in DPO dataset were filtered out. Check contents.")
|
||||
return None
|
||||
if after_len < before_len:
|
||||
logger.info(f"Filtered out {before_len - after_len} malformed rows; "
|
||||
f"{after_len} remain.")
|
||||
logger.info(f"DPO dataset ready with {after_len} samples.")
|
||||
|
||||
elif mode == "ftpo":
|
||||
if config.get("finetune_ftpo_dataset"):
|
||||
dataset_path = Path(config["finetune_ftpo_dataset"])
|
||||
else:
|
||||
ftpo_files = sorted(experiment_run_dir.glob("iter_*_ftpo_pairs.jsonl"))
|
||||
if not ftpo_files:
|
||||
logger.error("No ftpo files found for ftpo.")
|
||||
return None
|
||||
dataset_path = ftpo_files[-1]
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# FTPO dataset with dual regularisation + built-in size cap
|
||||
# ------------------------------------------------------------------
|
||||
dpo_dataset_hf = load_ftpo_multi_dataset(
|
||||
dataset_path,
|
||||
tokenizer,
|
||||
experiment_run_dir = experiment_run_dir,
|
||||
max_seq_len = max_seq_length,
|
||||
# balance *rejected* tokens
|
||||
rejected_reg_strength = config.get("ftpo_sample_rejected_regularisation_strength", 0.8),
|
||||
# balance *chosen* tokens
|
||||
chosen_reg_strength = config.get("ftpo_sample_chosen_regularisation_strength", 0.2),
|
||||
# hard floor on |chosen|
|
||||
min_chosen_tokens = config.get("ftpo_sample_min_chosen_tokens", 3),
|
||||
# overall training-set cap (used for per-token quotas too)
|
||||
max_train_examples = config.get("finetune_max_train_examples"),
|
||||
)
|
||||
|
||||
# loader already returns a shuffled dataset; an extra shuffle is fine but optional
|
||||
dpo_dataset_hf = dpo_dataset_hf.shuffle(seed=config.get("finetune_shuffle_seed", 3407))
|
||||
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────────────────────
|
||||
# [DEBUG] Inspect last-5 prompt tokens + chosen / rejected token
|
||||
# –– prints up to 50 ftpo examples for a quick sanity check.
|
||||
# –– gated by new config flag `finetune_debug_ftpo_tokens`.
|
||||
# ──────────────────────────────────────────────────────────────
|
||||
if False:
|
||||
sample_n = min(50, len(dpo_dataset_hf))
|
||||
print(f"\n🔎 ftpo debug: showing {sample_n} examples "
|
||||
"(last-5 prompt tokens, chosen ▸ rejected)\n")
|
||||
for i, ex in enumerate(dpo_dataset_hf.select(range(sample_n))):
|
||||
tail_prompt = tokenizer.convert_ids_to_tokens(ex["prompt_ids"][-5:])
|
||||
chosen_tok = tokenizer.convert_ids_to_tokens([ex["chosen_ids"][0]])[0]
|
||||
rejected_tok = tokenizer.convert_ids_to_tokens([ex["rejected_token_id"]])[0]
|
||||
tail_str = " ".join(tail_prompt)
|
||||
print(f"{i:03d}: … {tail_str} → {chosen_tok} ▸ {rejected_tok}")
|
||||
print("\n—— end ftpo debug ——\n")
|
||||
|
||||
elif mode == "dpo_final_token":
|
||||
# ------------------------------------------------------------
|
||||
# 1. Build the raw dataset **exactly** the same way FTPO does
|
||||
# ------------------------------------------------------------
|
||||
if config.get("finetune_ftpo_dataset"):
|
||||
dataset_path = Path(config["finetune_ftpo_dataset"])
|
||||
else:
|
||||
ftpo_files = sorted(experiment_run_dir.glob("iter_*_ftpo_pairs.jsonl"))
|
||||
if not ftpo_files:
|
||||
logger.error("No ftpo files found for dpo_final_token.")
|
||||
return None
|
||||
dataset_path = ftpo_files[-1]
|
||||
|
||||
ftpo_ds = load_ftpo_multi_dataset(
|
||||
dataset_path,
|
||||
tokenizer,
|
||||
experiment_run_dir = experiment_run_dir,
|
||||
max_seq_len = max_seq_length,
|
||||
rejected_reg_strength = config.get("ftpo_sample_rejected_regularisation_strength", 0.8),
|
||||
chosen_reg_strength = config.get("ftpo_sample_chosen_regularisation_strength", 0.2),
|
||||
min_chosen_tokens = config.get("ftpo_sample_min_chosen_tokens", 3),
|
||||
max_train_examples = config.get("finetune_max_train_examples"),
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# 2. Convert each row into a *single-token* DPO pair
|
||||
# ------------------------------------------------------------
|
||||
pairs = []
|
||||
pad_id = tokenizer.pad_token_id
|
||||
|
||||
for ex in ftpo_ds:
|
||||
# ––– recover the left-padded prompt as text –––
|
||||
prompt_ids = [tid for tid in ex["prompt_ids"] if tid != pad_id]
|
||||
prompt_txt = tokenizer.decode(prompt_ids, skip_special_tokens=False)
|
||||
|
||||
# ––– single-token continuations –––
|
||||
chosen_txt = tokenizer.decode(
|
||||
[ex["chosen_ids"][0]], skip_special_tokens=False
|
||||
)
|
||||
rejected_txt = tokenizer.decode(
|
||||
[ex["rejected_token_id"]], skip_special_tokens=False
|
||||
)
|
||||
|
||||
pairs.append(
|
||||
{
|
||||
"prompt": prompt_txt,
|
||||
"chosen": chosen_txt, # continuation only!
|
||||
"rejected": rejected_txt, # continuation only!
|
||||
}
|
||||
)
|
||||
|
||||
dpo_dataset_hf = Dataset.from_list(pairs)
|
||||
|
||||
# ── DEBUG: inspect a few prompt / chosen / rejected triples ──────────────
|
||||
def _show_examples(ds, n=3):
|
||||
for i, ex in enumerate(ds.select(range(n))):
|
||||
print(f"\n── example {i} ──")
|
||||
print("PROMPT:\n", ex["prompt"])
|
||||
print("CHOSEN:\n", ex["chosen"])
|
||||
print("REJECTED:\n", ex["rejected"])
|
||||
print("-" * 40)
|
||||
|
||||
_show_examples(dpo_dataset_hf, n=3)
|
||||
|
||||
# ------------------------------------------------------------
|
||||
# 3. Apply the *same* length filter & book-keeping as vanilla DPO
|
||||
# ------------------------------------------------------------
|
||||
def _within_len(example):
|
||||
p = tokenizer(example["prompt"], add_special_tokens=False).input_ids
|
||||
c = tokenizer(example["chosen"], add_special_tokens=False).input_ids
|
||||
r = tokenizer(example["rejected"],add_special_tokens=False).input_ids
|
||||
return len(p) + len(c) <= max_seq_length and len(p) + len(r) <= max_seq_length
|
||||
|
||||
before = len(dpo_dataset_hf)
|
||||
dpo_dataset_hf = dpo_dataset_hf.filter(_within_len)
|
||||
after = len(dpo_dataset_hf)
|
||||
logger.info(f"dpo_final_token length filter: kept {after}/{before} examples "
|
||||
f"(max_seq_len = {max_seq_length})")
|
||||
|
||||
if after == 0:
|
||||
raise ValueError("every sample exceeded finetune_max_seq_length")
|
||||
|
||||
max_train = config.get("finetune_max_train_examples")
|
||||
if isinstance(max_train, int) and max_train > 0 and len(dpo_dataset_hf) > max_train:
|
||||
dpo_dataset_hf = dpo_dataset_hf.select(range(max_train))
|
||||
logger.info(f"Capped training dataset to {max_train} examples.")
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────
|
||||
# ORPO — single-token pairs (prompt, chosen, rejected)
|
||||
# Mode value: "orpo_final_token"
|
||||
# ─────────────────────────────────────────────────────────────────────
|
||||
elif mode == "orpo_final_token":
|
||||
# 1) Construct the FTPO dataset exactly as in the ftpo branch
|
||||
if config.get("finetune_ftpo_dataset"):
|
||||
dataset_path = Path(config["finetune_ftpo_dataset"])
|
||||
else:
|
||||
ftpo_files = sorted(experiment_run_dir.glob("iter_*_ftpo_pairs.jsonl"))
|
||||
if not ftpo_files:
|
||||
logger.error("No ftpo files found for orpo_final_token.")
|
||||
return None
|
||||
dataset_path = ftpo_files[-1]
|
||||
|
||||
ftpo_ds = load_ftpo_multi_dataset(
|
||||
dataset_path,
|
||||
tokenizer,
|
||||
experiment_run_dir = experiment_run_dir,
|
||||
max_seq_len = max_seq_length,
|
||||
rejected_reg_strength = config.get("ftpo_sample_rejected_regularisation_strength", 0.8),
|
||||
chosen_reg_strength = config.get("ftpo_sample_chosen_regularisation_strength", 0.2),
|
||||
min_chosen_tokens = config.get("ftpo_sample_min_chosen_tokens", 3),
|
||||
max_train_examples = config.get("finetune_max_train_examples"),
|
||||
)
|
||||
|
||||
# 2) Convert to (prompt, chosen, rejected) triples — one per row
|
||||
pairs = []
|
||||
pad_id = tokenizer.pad_token_id
|
||||
|
||||
for ex in ftpo_ds:
|
||||
prompt_ids = [tid for tid in ex["prompt_ids"] if tid != pad_id]
|
||||
prompt_txt = tokenizer.decode(prompt_ids, skip_special_tokens=False)
|
||||
|
||||
chosen_txt = tokenizer.decode([ex["chosen_ids"][0]], skip_special_tokens=False)
|
||||
rejected_txt = tokenizer.decode([ex["rejected_token_id"]], skip_special_tokens=False)
|
||||
|
||||
pairs.append({"prompt": prompt_txt,
|
||||
"chosen": chosen_txt,
|
||||
"rejected": rejected_txt})
|
||||
|
||||
dpo_dataset_hf = Dataset.from_list(pairs)
|
||||
|
||||
# 3) Length filter / shuffle / cap (reuse helper)
|
||||
def _within_len(ex):
|
||||
p = tokenizer(ex["prompt"], add_special_tokens=False).input_ids
|
||||
c = tokenizer(ex["chosen"], add_special_tokens=False).input_ids
|
||||
r = tokenizer(ex["rejected"], add_special_tokens=False).input_ids
|
||||
return len(p) + len(c) <= max_seq_length and len(p) + len(r) <= max_seq_length
|
||||
|
||||
before = len(dpo_dataset_hf)
|
||||
dpo_dataset_hf = dpo_dataset_hf.filter(_within_len)
|
||||
logger.info(f"orpo_final_token length filter: kept {len(dpo_dataset_hf)}/{before} samples")
|
||||
|
||||
dpo_dataset_hf = dpo_dataset_hf.shuffle(seed=config.get("finetune_shuffle_seed", 3407))
|
||||
max_train = config.get("finetune_max_train_examples")
|
||||
if isinstance(max_train, int) and max_train > 0 and len(dpo_dataset_hf) > max_train:
|
||||
dpo_dataset_hf = dpo_dataset_hf.select(range(max_train))
|
||||
logger.info(f"Capped training dataset to {max_train} examples.")
|
||||
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────
|
||||
# KTO — single-token, unpaired (prompt, completion, label)
|
||||
# Mode value: "kto_final_token"
|
||||
# ─────────────────────────────────────────────────────────────────────
|
||||
elif mode == "kto_final_token":
|
||||
# 1) build the FTPO dataset exactly as before … (unchanged)
|
||||
ftpo_ds = load_ftpo_multi_dataset(
|
||||
dataset_path,
|
||||
tokenizer,
|
||||
experiment_run_dir = experiment_run_dir,
|
||||
max_seq_len = max_seq_length,
|
||||
rejected_reg_strength = config.get("ftpo_sample_rejected_regularisation_strength", 0.8),
|
||||
chosen_reg_strength = config.get("ftpo_sample_chosen_regularisation_strength", 0.2),
|
||||
min_chosen_tokens = config.get("ftpo_sample_min_chosen_tokens", 3),
|
||||
max_train_examples = config.get("finetune_max_train_examples"),
|
||||
)
|
||||
|
||||
# 2) ONE positive + ONE negative row per prompt ──────────────────────
|
||||
rows, pad_id = [], tokenizer.pad_token_id
|
||||
for ex in ftpo_ds:
|
||||
prompt_ids = [tid for tid in ex["prompt_ids"] if tid != pad_id]
|
||||
prompt_txt = tokenizer.decode(prompt_ids, skip_special_tokens=False)
|
||||
|
||||
if not ex["chosen_ids"]:
|
||||
continue # skip degenerate prompt
|
||||
|
||||
# positive (first chosen id)
|
||||
pos_txt = tokenizer.decode([ex["chosen_ids"][0]], skip_special_tokens=False)
|
||||
rows.append({"prompt": prompt_txt, "completion": pos_txt, "label": True})
|
||||
|
||||
# negative
|
||||
neg_txt = tokenizer.decode([ex["rejected_token_id"]], skip_special_tokens=False)
|
||||
rows.append({"prompt": prompt_txt, "completion": neg_txt, "label": False})
|
||||
|
||||
dpo_dataset_hf = Dataset.from_list(rows)
|
||||
|
||||
# 3) length filter ───────────────────────────────────────────────────
|
||||
def _within_len(ex):
|
||||
p = tokenizer(ex["prompt"], add_special_tokens=False).input_ids
|
||||
c = tokenizer(ex["completion"], add_special_tokens=False).input_ids
|
||||
return len(p) + len(c) <= max_seq_length
|
||||
|
||||
before = len(dpo_dataset_hf)
|
||||
dpo_dataset_hf = dpo_dataset_hf.filter(_within_len)
|
||||
logger.info(f"kto_final_token length filter: kept {len(dpo_dataset_hf)}/{before} samples")
|
||||
|
||||
# 4) cap first, then shuffle ─────────────────────────────────────────
|
||||
max_train = config.get("finetune_max_train_examples")
|
||||
if isinstance(max_train, int) and max_train > 0 and len(dpo_dataset_hf) > max_train:
|
||||
dpo_dataset_hf = dpo_dataset_hf.select(range(max_train))
|
||||
logger.info(f"Capped training dataset to {max_train} examples.")
|
||||
|
||||
#dpo_dataset_hf = dpo_dataset_hf.shuffle(seed=config.get("finetune_shuffle_seed", 3407))
|
||||
|
||||
else:
|
||||
logger.error(f"Unknown finetune_mode '{mode}'. Use 'dpo' or 'ftpo'.")
|
||||
return None
|
||||
|
||||
return dpo_dataset_hf
|
||||
Reference in New Issue
Block a user