lower num proc + better logging + self-gen now trainable

This commit is contained in:
51616 2025-07-03 15:46:06 +00:00
parent f803fd31b3
commit 21d8beb65d
4 changed files with 104 additions and 50 deletions

View file

@ -269,6 +269,7 @@ def main():
_get_tokenized_dataset = partial(
get_tokenized_dataset,
max_qas_len=ctx_args.max_qas_len,
base_model_max_len=model.base_model.config.max_position_embeddings,
tokenizer=tokenizer,
tokenizer_kwargs=tokenizer_kwargs,

View file

@ -230,12 +230,12 @@ def pack_batch(
f"# Packed samples: {len(idx_pairs)}\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\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\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)}"
f"Min: {np.min(packed_ctx_lens_arr)}, Max: {np.max(packed_ctx_lens_arr)}\n\n"
)
return packed_batch
@ -247,8 +247,6 @@ if __name__ == "__main__":
from ctx_to_lora.data.processing import get_tokenized_dataset
from ctx_to_lora.model_loading import get_model_and_tokenizer
disable_caching()
setup_logging("tmp/packing_debug.log", debug=True)
logger.info("Starting packing script...")
model, tokenizer = get_model_and_tokenizer(
@ -278,6 +276,8 @@ if __name__ == "__main__":
)
# ds.set_format("torch")
print(ds)
disable_caching()
packed_ds = ds.map(
pack_batch,
fn_kwargs={
@ -287,7 +287,7 @@ if __name__ == "__main__":
batched=True,
batch_size=100_000,
remove_columns=ds.column_names,
num_proc=8,
num_proc=4,
)
print(packed_ds)

View file

@ -312,10 +312,27 @@ def get_preprocessing_fn(
def get_ds_kwargs(ds_name: str, split: str) -> dict[str, Any]:
# custom logic for slicing iterable datasets
take, skip = None, None
if ("[" in split) and split.endswith("]"):
split, slice = split.split("[")
slice = slice.strip("]")
skip = slice.split(":")[0]
take = slice.split(":")[1]
if ds_name.startswith("self_gen/"):
# e.g., "self_gen/google/gemma-2-2b-it/pwc"
base_model_name = "/".join(ds_name.split("/")[1:-1])
base_ds = ds_name.split("/")[-1]
if ("[" in split) and split.endswith("]"):
kwargs["split"], slice = split.split("[")
slice = slice.strip("]")
skip = slice.split(":")[0]
if skip:
kwargs["skip"] = int(skip)
take = slice.split(":")[1]
if take:
kwargs["take"] = int(take)
files = glob(
f"{SELF_GEN_DATA_DIR}/{base_model_name}/{base_ds}/{split}/*.parquet"
)
@ -334,18 +351,11 @@ def get_ds_kwargs(ds_name: str, split: str) -> dict[str, Any]:
else:
kwargs = DS_KWARGS[ds_name][split]
# # custom logic for slicing iterable datasets
# take, skip = None, None
# kw_split = kwargs["split"]
# if ("[" in kw_split) and kw_split.endswith("]"):
# kwargs["split"], slice = kw_split.split("[")
# slice = slice.strip("]")
# skip = slice.split(":")[0]
# if skip:
# kwargs["skip"] = int(skip)
# take = slice.split(":")[1]
# if take:
# kwargs["take"] = int(take)
if skip:
kwargs["skip"] = int(skip)
if take:
kwargs["take"] = int(take)
return kwargs
@ -517,7 +527,7 @@ def load_and_process_dataset(
trust_remote_code=True,
streaming=load_as_stream,
)
if split == "train" and streaming and (not load_as_stream):
if "train" in split and streaming and (not load_as_stream):
# if the dataset is hosted on HF, we load it first then convert to iterable
ds = ds.to_iterable_dataset(num_shards=128)
if skip is not None:
@ -547,7 +557,7 @@ def load_and_process_dataset(
num_proc=num_proc,
)
if split == "train":
if "train" in split:
# TODO: drop samples or split into multiple contexts
# for multi-lora training
# ds = ds.filter(filter_long_samples, batched=True, num_proc=num_proc)
@ -636,13 +646,13 @@ def get_tokenized_dataset(
logger.debug(f"Dataset hash: {ds_hash}")
ds_path = f"{TRANSFORMED_DATA_DIR}/{ds_hash}"
if path.exists(ds_path) and split == "train" and is_caching_enabled():
if path.exists(ds_path) and ("train" in split) and is_caching_enabled():
# load the cached ds
logger.info(f"Loaded tokenized dataset from {ds_path}")
tokenized_ds = datasets.load_from_disk(ds_path)
return tokenized_ds
num_proc = None if streaming and split == "train" else 8
num_proc = None if streaming and ("train" in split) else 8
ds = load_and_process_dataset(
**load_and_process_kwargs,
num_proc=num_proc,
@ -656,7 +666,7 @@ def get_tokenized_dataset(
**tokenize_kwargs,
)
tokenized_ds = tokenized_ds.shuffle()
if split == "train" and is_caching_enabled():
if ("train" in split) and is_caching_enabled():
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)
@ -725,6 +735,7 @@ def construct_and_tokenize_ctx_qa(
)
# add input_ids, attention_mask, labels to "data"
os.environ["TOKENIZERS_PARALLELISM"] = "true"
logging.debug("Tokenizing inputs")
tokenized_ds = ds.map(
get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer),
batched=True,
@ -732,7 +743,7 @@ def construct_and_tokenize_ctx_qa(
# # num_proc=num_proc,
)
if split == "train":
if "train" in split:
# base model cant handle these samples naturally
# should skip
# tokenized_ds = tokenized_ds.filter(
@ -776,6 +787,7 @@ def construct_and_tokenize_ctx_qa(
if need_ctx_ids:
# tokenize the ctx_text to get ctx_ids and ctx_attn_mask
os.environ["TOKENIZERS_PARALLELISM"] = "true"
logging.debug("Tokenizing context")
tokenized_ds = tokenized_ds.map(
tokenize_ctx_text,
fn_kwargs={"tokenizer": ctx_tokenizer},
@ -783,7 +795,7 @@ def construct_and_tokenize_ctx_qa(
batch_size=100_000,
# # num_proc=num_proc,
)
if split == "train":
if "train" in split:
# TODO: do something if ctx length is longer than the ctx model length
# e.g., drop or split for multi-lora training
tokenized_ds = tokenized_ds.filter(
@ -793,19 +805,21 @@ def construct_and_tokenize_ctx_qa(
)
if not is_pretrain:
if split == "train":
if "train" in split:
logging.debug("Unpacking data")
tokenized_ds = tokenized_ds.map(
unpack_data,
num_proc=num_proc,
remove_columns=["data", "context"],
)
logging.debug(f"Split too long QAs with max length {max_qas_len}")
tokenized_ds = tokenized_ds.map(
split_too_long_qas,
fn_kwargs={"max_qas_len": max_qas_len},
batched=True,
batch_size=100_000,
num_proc=num_proc,
num_proc=4,
)
else:
cols_to_remove = ["data", "context"]
@ -997,54 +1011,94 @@ def split_too_long_qas(samples: dict[str, any], max_qas_len: int):
labels = samples["labels"]
ctx_ids = samples["ctx_ids"]
ctx_attn_mask = samples["ctx_attn_mask"]
# Pre-calculate total lengths to check if any splitting is needed
total_lengths = [sum(len(x) for x in seq) for seq in input_ids]
longest_old_qas_len = max(total_lengths) if total_lengths else 0
# Early exit if no splitting needed
if all(length <= max_qas_len for length in total_lengths):
logger.debug(f"Longest old qas len: {longest_old_qas_len}")
logger.debug(f"Longest new qas len: {longest_old_qas_len}")
return samples
out = {k: list() for k in samples}
longest_old_qas_len = 0
longest_new_qas_len = 0
for i in range(len(input_ids)):
tot_inp_len = sum([len(x) for x in input_ids[i]])
longest_old_qas_len = max(longest_old_qas_len, tot_inp_len)
n_skip = 0
# Helper function to add a batch efficiently
def add_batch(inp_ids_batch, attn_batch, labels_batch, ctx_id, ctx_attn):
out["input_ids"].append(inp_ids_batch)
out["attention_mask"].append(attn_batch)
out["labels"].append(labels_batch)
out["ctx_ids"].append(ctx_id)
out["ctx_attn_mask"].append(ctx_attn)
for i, tot_inp_len in enumerate(total_lengths):
if tot_inp_len <= max_qas_len:
# total length of the qas are shorter than max_qas_len
# no need to split
# No need to split - add entire sample
for k in samples:
out[k].append(samples[k][i])
continue
# Need to split this sample
current_ctx_id = ctx_ids[i]
current_ctx_attn = ctx_attn_mask[i]
new_qas_len = 0
new_input_ids = []
new_attn_mask = []
new_labels = []
for inp_ids, attn_mask, label in zip(
input_ids[i], attention_mask[i], labels[i]
):
if len(inp_ids) > max_qas_len:
# skip
inp_len = len(inp_ids)
if inp_len > max_qas_len:
# Skip individual sequences that are too long
n_skip += 1
continue
if new_qas_len + len(inp_ids) <= max_qas_len:
new_qas_len += len(inp_ids)
if new_qas_len + inp_len <= max_qas_len:
# Add to current batch
new_qas_len += inp_len
new_input_ids.append(inp_ids)
new_attn_mask.append(attn_mask)
new_labels.append(label)
else:
out["input_ids"].append(new_input_ids)
out["attention_mask"].append(new_attn_mask)
out["labels"].append(new_labels)
out["ctx_ids"].append(ctx_ids[i])
out["ctx_attn_mask"].append(ctx_attn_mask[i])
longest_new_qas_len = max(longest_new_qas_len, new_qas_len)
new_qas_len = len(inp_ids)
# Current batch is full, save it and start new batch
if new_input_ids: # Only add non-empty batches
add_batch(
new_input_ids,
new_attn_mask,
new_labels,
current_ctx_id,
current_ctx_attn,
)
longest_new_qas_len = max(longest_new_qas_len, new_qas_len)
# Start new batch with current sequence
new_qas_len = inp_len
new_input_ids = [inp_ids]
new_attn_mask = [attn_mask]
new_labels = [label]
# Add final batch if not empty
if new_input_ids:
add_batch(
new_input_ids,
new_attn_mask,
new_labels,
current_ctx_id,
current_ctx_attn,
)
longest_new_qas_len = max(longest_new_qas_len, new_qas_len)
out["input_ids"].append(new_input_ids)
out["attention_mask"].append(new_attn_mask)
out["labels"].append(new_labels)
out["ctx_ids"].append(ctx_ids[i])
out["ctx_attn_mask"].append(ctx_attn_mask[i])
logger.debug(f"Longest old qas len: {longest_old_qas_len}")
logger.debug(f"Longest new qas len: {longest_new_qas_len}")
if n_skip:
logger.warning(
f"Skipped {n_skip} QA pairs because they were too long (> {max_qas_len=} tokens)"
)
return out

View file

@ -74,7 +74,6 @@ def get_aggregator_config(
lora_r=lora_r,
per_rank_gen=per_rank_gen,
layer_to_layer_ctx_encoder=layer_to_layer_ctx_encoder,
n_ctx_model_layers=ctx_encoder_model_config.num_hidden_layers,
**vars(aggregator_args),
)