mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
lower num proc + better logging + self-gen now trainable
This commit is contained in:
parent
f803fd31b3
commit
21d8beb65d
4 changed files with 104 additions and 50 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue