pre-packing runnable

This commit is contained in:
51616 2025-06-10 15:41:18 +00:00
parent c8e14758ce
commit 76d2b154fd
6 changed files with 435 additions and 51 deletions

View file

@ -40,13 +40,14 @@ WANDB_MODE=disabled uv run intx_sft.py configs/pwc_hotpot_qa.yaml \
--gradient_accumulation_steps=1 --per_device_eval_batch_size=32 --exp_setup=hyper_lora \
--aggregator_type=perceiver \
--target_modules=down_proj \
--num_self_attends_per_block=4 --num_latent_factor=2 \
--num_self_attends_per_block=8 --num_latent_factor=2 \
--lora_r=8 \
--eval_steps=1000 --save_steps=1000 --learning_rate=4e-5 --lora_dropout=0.0 \
--neftune_noise_alpha=5 --use_light_weight_lora=False \
--load_best_model_at_end=True --metric_for_best_model=pwc_loss --add_negative_prompt=False \
--load_best_model_at_end=True --metric_for_best_model=eval_pwc_loss --add_negative_prompt=False \
--add_repeat_prompt=False \
--use_sequence_packing=True --per_rank_gen=True \
--use_sequence_packing=True --max_packed_inp_len=20000 --max_packed_ctx_len=40000 \
--per_rank_gen=True \
--per_layer_processing=True \
--gen_lora_l1_reg_coef=0.1 \
```

View file

@ -1,8 +1,9 @@
import contextlib
import logging
import os
from copy import deepcopy
from functools import partial
from math import ceil, isclose
from math import isclose
import numpy as np
import torch
@ -30,8 +31,11 @@ from ctx_to_lora.configs import (
ModelArguments,
TrainingArguments,
)
from ctx_to_lora.data.collator import train_collator, train_packed_collator
from ctx_to_lora.data.processing import get_tokenized_dataset
from ctx_to_lora.data.collator import ( # train_packed_collator,; DefaultDataCollator,
flatten_if_not_packed,
train_collator,
)
from ctx_to_lora.data.processing import get_tokenized_dataset, pack
from ctx_to_lora.metrics import (
Evaluator,
compute_metrics,
@ -269,24 +273,33 @@ def main():
set_format=None if ctx_args.use_sequence_packing else "pt",
# streaming=data_args.streaming,
)
tokenized_ds = {"train": dict(), "validation": dict(), "test": dict()}
tokenized_ds = {"train": dict(), "validation": dict()}
for split, ds_names in zip(
["train", "validation", "test"],
[data_args.train_ds_names, data_args.val_ds_names, data_args.test_ds_names],
["train", "validation"],
[data_args.train_ds_names, data_args.val_ds_names],
):
if not ds_names:
continue
streaming = (split == "train") and data_args.streaming
for ds_name in ds_names:
ds = _get_tokenized_dataset(ds_name, split, streaming=streaming)
ctx_mgr = (
training_args.main_process_first()
if split == "train"
else contextlib.nullcontext()
)
with ctx_mgr:
# process and tokenize on the main process
# then other replicas can just load the cached dataset
# we dont save cache for validation ds
for ds_name in ds_names:
ds = _get_tokenized_dataset(ds_name, split, streaming=streaming)
base_name = os.path.basename(ds_name)
if ds_name.startswith("self_gen/"):
ds_name = "self_gen/" + base_name
else:
ds_name = base_name
base_name = os.path.basename(ds_name)
if ds_name.startswith("self_gen/"):
ds_name = "self_gen/" + base_name
else:
ds_name = base_name
tokenized_ds[split][ds_name] = ds
tokenized_ds[split][ds_name] = ds
train_ds = tokenized_ds["train"]
logging.info(f"train_ds: {train_ds}")
@ -328,27 +341,43 @@ def main():
# seed=training_args.seed,
# )
# else:
train_ds_len = [len(ds) for ds in train_ds.values()]
total_len = sum(train_ds_len)
max_steps = ceil(
total_len
* training_args.num_train_epochs
/ training_args.per_device_train_batch_size
/ training_args.gradient_accumulation_steps
/ training_args.world_size
)
training_args.max_steps = max_steps
# max_steps = ceil(
# total_len
# * training_args.num_train_epochs
# / training_args.per_device_train_batch_size
# / training_args.gradient_accumulation_steps
# / training_args.world_size
# )
# training_args.max_steps = max_steps
train_ds = interleave_datasets(
list(train_ds.values()),
probabilities=get_ds_prob(train_ds_len, total_len),
seed=training_args.seed,
stopping_strategy="all_exhausted",
)
if ctx_args.use_sequence_packing:
logging.info("Packing dataset")
train_ds = pack(
train_ds,
ctx_args.max_packed_inp_len,
ctx_args.max_packed_ctx_len,
max_packed_size=-1,
num_proc=8,
)
# TODO: add stats here
logging.info("Setting per_device_train_batch_size to 1")
training_args.per_device_train_batch_size = 1
logger.info(f"train_ds: {train_ds}")
logger.info(f"val_ds: {val_ds}")
collator = (
train_packed_collator
flatten_if_not_packed
# DefaultDataCollator(return_tensors="pt")
if ctx_args.use_sequence_packing
else partial(train_collator, tokenizer=tokenizer)
)

View file

@ -307,12 +307,16 @@ class CtxTrainingArguments:
default=False,
metadata={"help": "Whether to use sequence packing."},
)
per_device_train_max_batch_len: int | None = field(
default=2**12,
metadata={
"help": "Maximum batch length for training. Only used with multipack sampler."
},
max_packed_inp_len: int | None = field(
default=2**14,
metadata={"help": "Maximum packed input length for training."},
)
max_packed_ctx_len: int | None = field(
# forward pass of the ctx encoder is cheaper --> longer packed len
default=2**15,
metadata={"help": "Maximum packed context length for training."},
)
max_new_tokens: int | None = field(
default=2**10,
metadata={"help": "Maximum new tokens for generation-based evaluation."},
@ -355,6 +359,7 @@ class DataArguments:
default=None,
metadata={"help": "Training dataset names."},
)
streaming: bool = field(
default=False,
metadata={"help": "Whether to use streaming dataset for training."},

View file

@ -1,20 +1,44 @@
import numpy as np
import torch
from transformers.data import DataCollatorWithFlattening
from transformers.data import (
DataCollatorWithFlattening,
default_data_collator,
)
flattener = DataCollatorWithFlattening()
def train_packed_collator(inp_list):
# def train_packed_collator(inp_list):
# # no padding
# packed_inputs = flattener(inp_list, return_tensors="pt")
# if "ctx_ids" in inp_list[0]:
# ctx_ids = [{"input_ids": example["ctx_ids"]} for example in inp_list]
# packed_ctx = flattener(ctx_ids, return_tensors="pt")
# packed_inputs["ctx_ids"] = packed_ctx["input_ids"]
# packed_inputs["ctx_position_ids"] = packed_ctx["position_ids"]
# # for eval info
# if "ctx_ids_len" in inp_list[0]:
# packed_inputs["ctx_ids_len"] = [
# example["ctx_ids_len"] for example in inp_list
# ]
# return packed_inputs
def flatten_if_not_packed(inp_list):
# no padding
sample = inp_list[0]
if "position_ids" in sample:
return default_data_collator(inp_list, return_tensors="pt")
packed_inputs = flattener(inp_list, return_tensors="pt")
if "ctx_ids" in inp_list[0]:
if "ctx_ids" in sample:
ctx_ids = [{"input_ids": example["ctx_ids"]} for example in inp_list]
packed_ctx = flattener(ctx_ids, return_tensors="pt")
packed_inputs["ctx_ids"] = packed_ctx["input_ids"]
packed_inputs["ctx_position_ids"] = packed_ctx["position_ids"]
# for eval info
if "ctx_ids_len" in inp_list[0]:
if "ctx_ids_len" in sample:
packed_inputs["ctx_ids_len"] = [
example["ctx_ids_len"] for example in inp_list
]

View file

@ -0,0 +1,279 @@
# based on
# https://github.com/MeetKai/functionary/blob/aa3dbdd65f7e388f2386622606bdfeec95c2b863/functionary/train/packing/packed_dataset.py
import logging
import numpy as np
logger = logging.getLogger()
def pack_data_points_by_length(
lens: list[int],
ctx_lens: list[int],
max_packed_inp_len: int,
max_packed_ctx_len: int,
max_size: int = -1,
) -> list[int]:
"""Fully vectorized version using NumPy operations for maximum performance.
This is an alternative implementation that's more complex but potentially faster
for very large datasets.
"""
if not lens:
return []
len_arr = np.array(lens, dtype=np.int32)
ctx_len_arr = np.array(ctx_lens, dtype=np.int32)
n = len(len_arr)
if n == 1:
return (
[[0]]
if len_arr[0] <= max_packed_inp_len or ctx_len_arr[0] <= max_packed_ctx_len
else []
)
# Create cumulative sum arrays for efficient range queries
cumsum_inp_len = np.cumsum(len_arr)
cumsum_ctx_len = np.cumsum(ctx_len_arr)
boundaries = [0]
i = 0
while i < n:
# Find the maximum j such that sum(lens[i:j+1]) <= max_packed_inp_len
start_sum_inp = cumsum_inp_len[i - 1] if i > 0 else 0
valid_ends_inp = cumsum_inp_len[i:] - start_sum_inp <= max_packed_inp_len
start_sum_ctx = cumsum_ctx_len[i - 1] if i > 0 else 0
valid_ends_ctx = cumsum_ctx_len[i:] - start_sum_ctx <= max_packed_ctx_len
valid_ends = valid_ends_inp & valid_ends_ctx
# this should never happen?
# if not np.any(valid_ends):
# # Single item exceeds max_packed_inp_len, skip it
# i += 1
# continue
# Find the last valid index
max_valid_idx = i + np.where(valid_ends)[0][-1]
# Apply max_size constraint
if max_size != -1:
max_valid_idx = min(max_valid_idx, i + max_size - 1)
boundaries.append(max_valid_idx + 1)
i = max_valid_idx + 1
return boundaries
def pack_data_points_FA(
batch: dict[str, any],
) -> dict[str, np.ndarray]:
"""Optimized packing function using torch operations."""
if not batch:
raise ValueError("Batch is empty")
# Pre-allocate lists with known sizes
total_ctx_len = sum(len(x) for x in batch["ctx_ids"])
total_inp_len = sum(len(x) for x in batch["input_ids"])
ctx_ids = np.empty(total_ctx_len, dtype=np.long)
ctx_position_ids = np.empty(total_ctx_len, dtype=np.long)
input_ids = np.empty(total_inp_len, dtype=np.long)
position_ids = np.empty(total_inp_len, dtype=np.long)
labels = np.empty(total_inp_len, dtype=np.long)
offset = ctx_offset = 0
for ctx_ids_b, input_ids_b, labels_b in zip(
batch["ctx_ids"], batch["input_ids"], batch["labels"]
):
ctx_len = len(ctx_ids_b)
inp_len = len(input_ids_b)
inp_start, inp_end = offset, offset + inp_len
input_ids[inp_start:inp_end] = input_ids_b
position_ids[inp_start:inp_end] = np.arange(inp_len, dtype=np.long)
labels[inp_start:inp_end] = labels_b
offset += inp_len
ctx_start, ctx_end = ctx_offset, ctx_offset + ctx_len
ctx_ids[ctx_start:ctx_end] = ctx_ids_b
ctx_position_ids[ctx_start:ctx_end] = np.arange(ctx_len, dtype=np.long)
ctx_offset += ctx_len
return {
"ctx_ids": ctx_ids,
"ctx_position_ids": ctx_position_ids,
"input_ids": input_ids,
"position_ids": position_ids,
"labels": labels,
}
def pack_batch(
batch: dict[str, any],
max_packed_inp_len: int,
max_packed_ctx_len: int,
max_packed_size: int = -1,
) -> dict[str, any]:
# Extract lengths
# we do not pad so we can just take the length of the tokens
inp_lens = [len(x) for x in batch["input_ids"]]
inp_count = len(inp_lens)
# total_original_tokens = sum(inp_lens)
if "ctx_ids" not in batch:
raise ValueError("Batch must contain 'ctx_ids' and 'labels' keys")
# we do not pad so we can just take the length of the tokens
ctx_lens = [len(x) for x in batch["ctx_ids"]]
# Group indices
boundaries = pack_data_points_by_length(
inp_lens,
ctx_lens,
max_packed_inp_len,
max_packed_ctx_len,
max_packed_size,
)
# Pack groups
n_samples = len(boundaries) - 1
packed_batch = {
"ctx_ids": [],
"ctx_position_ids": [],
"input_ids": [],
"position_ids": [],
"labels": [],
}
packing_efficiency_ratios = []
ctx_packing_efficiency_ratios = []
for i in range(len(boundaries) - 1):
start_idx = boundaries[i]
end_idx = boundaries[i + 1]
group_items = {
"ctx_ids": batch["ctx_ids"][start_idx:end_idx],
"input_ids": batch["input_ids"][start_idx:end_idx],
"labels": batch["labels"][start_idx:end_idx],
}
packed_item = pack_data_points_FA(group_items)
packed_batch["ctx_ids"].append(packed_item["ctx_ids"])
packed_batch["ctx_position_ids"].append(packed_item["ctx_position_ids"])
packed_batch["input_ids"].append(packed_item["input_ids"])
packed_batch["position_ids"].append(packed_item["position_ids"])
packed_batch["labels"].append(packed_item["labels"])
# # Track packing efficiency
# group_inp_tokens = sum(inp_lens[start_idx:end_idx])
# group_ctx_tokens = sum(ctx_lens[start_idx:end_idx])
# if max_packed_inp_len > 0:
# inp_efficiency = group_inp_tokens / max_packed_inp_len
# packing_efficiency_ratios.append(inp_efficiency)
# if max_packed_ctx_len > 0:
# ctx_efficiency = group_ctx_tokens / max_packed_ctx_len
# ctx_packing_efficiency_ratios.append(ctx_efficiency)
# # Calculate total original tokens
# total_original_inp_tokens = sum(inp_lens)
# total_original_ctx_tokens = sum(ctx_lens)
# total_original_tokens = total_original_inp_tokens + total_original_ctx_tokens
# # Calculate length statistics
# packed_inp_lens_arr = np.array([len(x) for x in packed_batch["input_ids"]])
# packed_ctx_lens_arr = np.array([len(x) for x in packed_batch["ctx_ids"]])
# # Log performance statistics
# compression_ratio = inp_count / n_samples if n_samples > 0 else 0
# avg_inp_packing_efficiency = (
# np.mean(packing_efficiency_ratios) if packing_efficiency_ratios else 0
# )
# avg_ctx_packing_efficiency = (
# np.mean(ctx_packing_efficiency_ratios) if ctx_packing_efficiency_ratios else 0
# )
# print(
# f"Packing stats - Original samples: {inp_count}\n"
# f"Packed samples: {n_samples}\n"
# f"Compression ratio (sample-wise): {compression_ratio:.2f}x\n\n"
# f"Total inp tokens: {total_original_inp_tokens}\n"
# f"Total ctx tokens: {total_original_ctx_tokens}\n"
# f"Avg inp packing efficiency: {avg_inp_packing_efficiency:.3f}\n"
# f"Avg ctx packing efficiency: {avg_ctx_packing_efficiency:.3f}\n\n"
# f"Input IDs length stats:\n"
# f" Avg: {np.mean(packed_inp_lens_arr):.1f}, Std: {np.std(packed_inp_lens_arr):.1f}, "
# f"Min: {np.min(packed_inp_lens_arr)}, Max: {np.max(packed_inp_lens_arr)}\n"
# f"Context IDs length stats:\n"
# f" Avg: {np.mean(packed_ctx_lens_arr):.1f}, Std: {np.std(packed_ctx_lens_arr):.1f}, "
# f"Min: {np.min(packed_ctx_lens_arr)}, Max: {np.max(packed_ctx_lens_arr)}"
# )
return packed_batch
if __name__ == "__main__":
from ctx_to_lora.data.processing import get_tokenized_dataset
from ctx_to_lora.model_loading import get_model_and_tokenizer
model, tokenizer = get_model_and_tokenizer(
"google/gemma-2-2b-it",
train=True,
requires_grad=False,
)
base_model_max_len = model.base_model.config.max_position_embeddings
tokenizer_kwargs = {"max_length": base_model_max_len} # not used
ctx_tokenizer_kwargs = {"max_length": base_model_max_len} # not used for now
ds = get_tokenized_dataset(
ds_name="fw_qa_3_mini_pretrain",
split="train",
base_model_max_len=model.base_model.config.max_position_embeddings,
tokenizer=tokenizer,
tokenizer_kwargs={},
ctx_model_max_len=model.base_model.config.max_position_embeddings,
ctx_tokenizer=tokenizer,
ctx_tokenizer_kwargs={},
add_ctx_to_chat=False,
add_repeat_prompt=False,
repeat_prob=0.0,
add_negative_prompt=False,
use_kl_loss=False,
)
# ds.set_format("torch")
print(ds)
packed_ds = ds.map(
pack_batch,
fn_kwargs={
"max_packed_inp_len": base_model_max_len * 2,
"max_packed_ctx_len": base_model_max_len * 4,
},
batched=True,
batch_size=1000,
remove_columns=ds.column_names,
num_proc=8,
)
print(packed_ds)
# packed_ds = PackedDataset(
# dataset=ds,
# tokenizer=tokenizer,
# max_input_length=base_model_max_len,
# max_packed_length=base_model_max_len * 2,
# # max_packed_size=3, # maximum number of data points being packed
# )
# packed_ds.stat()
# packed_ds = Dataset.from_list(packed_ds)
# packed_ds.set_format("torch")
# print(packed_ds)
# print(ds[0])
# print(packed_ds[0])
orig_seq = ds[0]["ctx_ids"]
packed_seq = packed_ds[0]["ctx_ids"][: len(orig_seq)]
print(orig_seq, packed_seq)
print(orig_seq == packed_seq)

View file

@ -9,7 +9,7 @@ from typing import Any
import datasets
import numpy as np
from datasets import load_dataset
from datasets import Dataset, load_dataset
from transformers import PreTrainedTokenizerBase
from ctx_to_lora.data.definitions import (
@ -21,6 +21,7 @@ from ctx_to_lora.data.definitions import (
SELF_GEN_DATA_DIR,
TRANSFORMED_DATA_DIR,
)
from ctx_to_lora.data.packing import pack_batch
from ctx_to_lora.utils import TRAINING_TASK
logger = logging.getLogger()
@ -294,7 +295,7 @@ def validate_columns(tokenized_ds):
def len_filter(sample, max_length: int, keys: list[str]):
m = [len(sample[k]) < max_length for k in keys]
m = [len(sample[k]) <= max_length for k in keys]
return sum(m) == len(keys)
@ -459,13 +460,13 @@ def load_and_process_dataset(
ds = ds.map(
get_preprocessing_fn(ds_name, is_eval, is_pretrain),
remove_columns=cols_to_remove,
# num_proc=num_proc,
num_proc=num_proc,
)
# ds = ds.remove_columns(cols_to_remove)
ds = ds.filter(
filter_none,
batched=False,
# num_proc=num_proc,
num_proc=num_proc,
)
if split == "train":
@ -573,7 +574,8 @@ def get_tokenized_dataset(
)
if split == "train":
tokenized_ds.save_to_disk(ds_path, num_proc=num_proc)
# force reload from disk for fingerprint consistency
tokenized_ds = datasets.load_from_disk(ds_path)
return tokenized_ds
@ -598,16 +600,17 @@ def construct_and_tokenize_ctx_qa(
if is_paraphrased:
ds = ds.map(
build_paraphrase_pretrain,
batched=True,
batch_size=100_000,
# batched=True,
# batch_size=100_000,
remove_columns=["variation"],
num_proc=num_proc,
)
else:
ds = ds.map(
build_intx_pretrain,
# batched=True,
# batch_size=100_000,
# # num_proc=num_proc,
num_proc=num_proc,
)
tokenized_ds = ds.map(
tokenize_pretrain,
@ -628,7 +631,7 @@ def construct_and_tokenize_ctx_qa(
ds = ds.map(
convert_ctx_prompt_response_to_messages,
fn_kwargs={"add_ctx_to_chat": add_ctx_to_chat},
# # num_proc=num_proc,
num_proc=num_proc,
)
# add "chat" field
tokenized_ds = ds.map(
@ -671,7 +674,7 @@ def construct_and_tokenize_ctx_qa(
},
# batched=False,
# batch_size=0,
# num_proc=num_proc,
num_proc=num_proc,
)
# tokenized_ds = len_filter(
# tokenized_ds, base_model_max_len, ["input_ids", "labels"]
@ -687,14 +690,14 @@ def construct_and_tokenize_ctx_qa(
},
# batched=True,
# batch_size=100_000,
# # num_proc=num_proc,
# num_proc=num_proc,
)
tokenized_ds = tokenized_ds.map(
add_length_info,
fn_kwargs={"columns": ["input_ids"]},
# batched=True,
# batch_size=100_000,
# # num_proc=num_proc,
# num_proc=num_proc,
)
# for use_kl_loss, we need "chat_ids" and "chat_attn_mask"
@ -736,7 +739,7 @@ def construct_and_tokenize_ctx_qa(
fn_kwargs={"max_length": ctx_model_max_len, "keys": ["ctx_ids"]},
# batched=False,
# batch_size=0,
# # num_proc=num_proc,
num_proc=num_proc,
)
# tokenized_ds = len_filter(tokenized_ds, ctx_model_max_len, ["ctx_ids"])
@ -752,14 +755,14 @@ def construct_and_tokenize_ctx_qa(
},
# batched=True,
# batch_size=100_000,
# # num_proc=num_proc,
# num_proc=num_proc,
)
tokenized_ds = tokenized_ds.map(
add_length_info,
fn_kwargs={"columns": ["ctx_ids"]},
# batched=True,
# batch_size=100_000,
# # num_proc=num_proc,
# num_proc=num_proc,
)
if is_pretrain:
@ -858,6 +861,7 @@ def get_sft_prompt_formatting_fn(
for tok_ids, masks in zip(tokens["input_ids"], tokens["assistant_masks"]):
o = [id_ if mask else IGNORE_INDEX for id_, mask in zip(tok_ids, masks)]
labels.append(o)
del tokens["assistant_masks"]
tokens["labels"] = labels
return tokens
@ -1059,7 +1063,6 @@ def tokenize_pretrain(
truncation=False,
**(tokenizer_kwargs or {}),
)
# if tokenizer.bos_token_id is not None:
# for i in range(len(tokens["input_ids"])):
# # add bos
@ -1137,6 +1140,49 @@ def tokenize_ctx_text(
return dict(ctx_ids=ctx_ids, ctx_attn_mask=ctx_attn_mask)
def pack(
ds: Dataset,
max_packed_inp_len: int,
max_packed_ctx_len: int,
max_packed_size: int,
num_proc: int = 0,
):
# TODO: packing has to happen after concat'ing all the ds
# might have to do this on the fly...
# or have a giant cached file for the already concat'd + packed ds
kwargs = dict(
max_packed_inp_len=max_packed_inp_len,
max_packed_ctx_len=max_packed_ctx_len,
max_packed_size=max_packed_size,
)
ds_hash = hashlib.sha256(
(ds._fingerprint + json.dumps(kwargs)).encode()
).hexdigest()
ds_path = f"{TRANSFORMED_DATA_DIR}/{ds_hash}"
logger.info(
f"Packing dataset {ds_hash} with max_packed_inp_len={max_packed_inp_len} and max_packed_ctx_len={max_packed_ctx_len}"
)
if path.exists(ds_path):
logger.info(f"Loading a cached packed dataset for {ds_path}")
return datasets.load_from_disk(ds_path)
else:
ds = ds.map(
pack_batch,
fn_kwargs={
"max_packed_inp_len": max_packed_inp_len,
"max_packed_ctx_len": max_packed_ctx_len,
"max_packed_size": max_packed_size,
},
batched=True,
batch_size=1_000_000,
num_proc=num_proc,
remove_columns=ds.column_names,
)
ds.save_to_disk(ds_path, num_proc=num_proc)
return ds
if __name__ == "__main__":
from transformers import AutoTokenizer