diff --git a/intx_sft.py b/intx_sft.py index a4575f1..25890db 100755 --- a/intx_sft.py +++ b/intx_sft.py @@ -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, diff --git a/src/ctx_to_lora/data/packing.py b/src/ctx_to_lora/data/packing.py index 07870ba..11aa612 100644 --- a/src/ctx_to_lora/data/packing.py +++ b/src/ctx_to_lora/data/packing.py @@ -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) diff --git a/src/ctx_to_lora/data/processing.py b/src/ctx_to_lora/data/processing.py index 2284c7e..9f22f35 100644 --- a/src/ctx_to_lora/data/processing.py +++ b/src/ctx_to_lora/data/processing.py @@ -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 diff --git a/src/ctx_to_lora/modeling/aggregator.py b/src/ctx_to_lora/modeling/aggregator.py index de51670..c116f3a 100644 --- a/src/ctx_to_lora/modeling/aggregator.py +++ b/src/ctx_to_lora/modeling/aggregator.py @@ -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), )