| |
| |
|
|
| import json |
| import itertools |
| import random |
| from dataclasses import dataclass |
| from typing import Optional |
|
|
| import torch |
| import torch.distributed as dist |
| from datasets import Dataset |
| from transformers import PreTrainedTokenizerBase |
| from transformers.data.data_collator import pad_without_fast_tokenizer_warning |
|
|
|
|
| @dataclass |
| class MyCollator: |
|
|
| tokenizer: PreTrainedTokenizerBase |
| latent_id: Optional[int] = None |
| label_pad_token_id: Optional[int] = -100 |
|
|
| def __call__(self, features, return_tensors=None): |
|
|
| assert self.tokenizer.padding_side == "right" |
|
|
| """ |
| Pad the batch like this to maximize the reuse of kv cache. |
| E.g., |
| |
| xxxxxxxxxx<latent><latent>xxxxx-- |
| -----xxxxx<latent>xxxxxxxx------- |
| ---xxxxxxx<latent><latent>xxxxxxx |
| |
| |
| ("x" is word token, "-" is pad token) |
| """ |
| |
| |
| earliest_latent = [ |
| feature["input_ids"].index(self.latent_id) |
| for feature in features |
| if self.latent_id in feature["input_ids"] |
| ] |
|
|
| if len(earliest_latent) > 0: |
| latest_earliest_latent = max(earliest_latent) |
| for feature in features: |
| if self.latent_id in feature["input_ids"]: |
| n_tok_pad = latest_earliest_latent - feature["input_ids"].index( |
| self.latent_id |
| ) |
| else: |
| n_tok_pad = 0 |
| feature["position_ids"] = [0] * n_tok_pad + list( |
| range(len(feature["input_ids"])) |
| ) |
| feature["input_ids"] = [ |
| self.tokenizer.pad_token_id |
| ] * n_tok_pad + feature["input_ids"] |
| if "labels" in feature: |
| feature["labels"] = [self.label_pad_token_id] * n_tok_pad + feature[ |
| "labels" |
| ] |
| feature["attention_mask"] = [0] * n_tok_pad + feature["attention_mask"] |
|
|
| return_tensors = "pt" |
|
|
| label_name = "label" if "label" in features[0].keys() else "labels" |
|
|
| non_label_position_features = [ |
| { |
| k: v |
| for k, v in feature.items() |
| if k != label_name and k != "position_ids" |
| } |
| for feature in features |
| ] |
|
|
| |
| batch = pad_without_fast_tokenizer_warning( |
| self.tokenizer, |
| non_label_position_features, |
| padding=True, |
| pad_to_multiple_of=None, |
| return_tensors=return_tensors, |
| ) |
|
|
| labels = ( |
| [feature[label_name] for feature in features] |
| if label_name in features[0].keys() |
| else None |
| ) |
| if labels is not None and all(label is None for label in labels): |
| labels = None |
| position_ids = ( |
| [feature["position_ids"] for feature in features] |
| if "position_ids" in features[0].keys() |
| else None |
| ) |
| |
|
|
| if labels is not None: |
| max_label_length = max(len(l) for l in labels) |
|
|
| batch["labels"] = [ |
| label + [self.label_pad_token_id] * (max_label_length - len(label)) |
| for label in labels |
| ] |
| batch["labels"] = torch.tensor(batch["labels"], dtype=torch.int64) |
|
|
| if position_ids is not None: |
| max_pos_length = max(len(l) for l in position_ids) |
|
|
| batch["position_ids"] = [ |
| position_id + [0] * (max_pos_length - len(position_id)) |
| for position_id in position_ids |
| ] |
| batch["position_ids"] = torch.tensor( |
| batch["position_ids"], dtype=torch.int64 |
| ) |
|
|
| return batch |
|
|
|
|
| def expand_data(data, k, max_steps, neg_sampling=False, stage_matched_q=False): |
| """Build (prompt, continuation) for hop budget k. |
| |
| Default (stage_matched_q=False): |
| [Q] always lists the final leaf + decoy; CE at mid stages is a random |
| depth-k frontier node (may not appear in [Q]). |
| |
| stage_matched_q=True (K's format): |
| At hop k <= L, [Q] lists two nodes at distance k — the correct node on |
| the path to the target leaf (steps[k-1]) and a decoy from |
| neg_neighbor_k[k]. CE asks the model to emit the correct one of the two. |
| Final stage (k = L+1) unchanged: [Q] = final leaf pair, CE = leaf. |
| """ |
| assert k <= max_steps + 1 |
| |
| |
| symbol_to_idx = {} |
| for i, s in enumerate(data['idx_to_symbol']): |
| symbol_to_idx[s] = i |
|
|
| def _q_candidates(cand_a, cand_b): |
| if random.random() < 0.5: |
| return str(cand_a) + " " + str(cand_b) |
| return str(cand_b) + " " + str(cand_a) |
|
|
| def get_prefix(data, hop_k=None): |
| random.shuffle(data['edges']) |
|
|
| question = "<eos> " + "|".join([f" {e[0]} {e[1]} " for e in data['edges']]).strip() + \ |
| " [Q] " |
|
|
| if stage_matched_q and hop_k is not None and 1 <= hop_k <= max_steps: |
| |
| correct = int(data["steps"][hop_k - 1]) |
| decoy = int(random.choice(data["neg_neighbor_k"][str(hop_k)])) |
| question += _q_candidates(correct, decoy) |
| else: |
| question += _q_candidates(data["target"], data["neg_target"]) |
|
|
| question += " [R] " + str(data['root']) |
| |
| return question |
|
|
|
|
| |
| if k <= max_steps: |
| |
| if neg_sampling: |
| if random.random() < 0.2: |
| question = get_prefix(data, hop_k=k) + " <|latent|>" * (k-1) + " [A] " |
| continuation = "<|no-answer|>" |
| return_data = (question, continuation) |
| |
| else: |
| if stage_matched_q: |
| n = int(data["steps"][k - 1]) |
| else: |
| n = random.choice(data["neighbor_k"][str(k)]) |
| question = get_prefix(data, hop_k=k) + " <|latent|>" * (k-1) + " " |
| continuation = str(n) |
| return_data = (question, continuation) |
| else: |
| if stage_matched_q: |
| n = int(data["steps"][k - 1]) |
| else: |
| n = random.choice(data["neighbor_k"][str(k)]) |
| question = get_prefix(data, hop_k=k) + " <|latent|>" * (k-1) + " " |
| continuation = str(n) |
| return_data = (question, continuation) |
|
|
|
|
| elif k == max_steps + 1: |
| if neg_sampling: |
| if random.random() < 0.2: |
| question = get_prefix(data) + " <|latent|>" * random.randint(0, max_steps - 1) + " [A] " |
| continuation = "<|no-answer|>" |
| return_data = (question, continuation) |
| else: |
| question = get_prefix(data) + " <|latent|>" * max_steps + " [A] " |
| continuation = str(data["target"]) |
| return_data = (question, continuation) |
| else: |
| question = get_prefix(data) + " <|latent|>" * max_steps + " [A] " |
| continuation = str(data["target"]) |
| return_data = (question, continuation) |
| |
| else: |
| raise ValueError(f"k is {k}, max_steps is {max_steps}") |
| |
| return return_data |
|
|
|
|
|
|
| def get_graph_latent_cot_dataset( |
| dataset_path, |
| scheduled_stage, |
| configs, |
| tokenizer, |
| ): |
| base_dataset = json.load(open(dataset_path)) |
| |
| if configs.debug: |
| base_dataset = base_dataset[:10000] |
| train_size = getattr(configs, "train_size", 0) |
| if train_size and train_size > 0: |
| base_dataset = base_dataset[:train_size] |
|
|
| def process_dataset(sample): |
|
|
| if ( |
| random.random() < configs.uniform_prob |
| ): |
| scheduled_stage_to_train = random.randint(0, min(scheduled_stage, len(sample["steps"]))) |
| else: |
| scheduled_stage_to_train = min(scheduled_stage, len(sample["steps"])) |
| |
|
|
| |
| |
|
|
| _smq = bool(getattr(configs, "stage_matched_q", False)) |
| expanded_data = expand_data( |
| sample, |
| scheduled_stage_to_train + 1, |
| len(sample["steps"]), |
| stage_matched_q=_smq, |
| ) |
| |
| |
| processed_samples = [] |
| for question, continuation in [expanded_data]: |
| question_tokenized = tokenizer.encode(question, add_special_tokens=False) |
| continuation_tokenized = tokenizer.encode(continuation, add_special_tokens=False) |
| |
| tokens = question_tokenized + continuation_tokenized |
| |
| processed_sample = { |
| "input_ids": tokens, |
| "labels": [-100] * len(question_tokenized) + continuation_tokenized, |
| "attention_mask": [1] * len(tokens), |
| "position_ids": list(range(len(tokens))), |
| } |
| processed_samples.append(processed_sample) |
| |
| return processed_samples |
|
|
| if torch.cuda.device_count() > 1: |
| if dist.get_rank() == 0: |
| |
| all_processed_samples = [] |
| for sample in base_dataset: |
| processed_samples = process_dataset(sample) |
| all_processed_samples.extend(processed_samples) |
| |
| processed_dataset = all_processed_samples |
| |
| random.shuffle(processed_dataset) |
| processed_dataset = [processed_dataset] |
| else: |
| processed_dataset = [None] |
| dist.broadcast_object_list(processed_dataset, src=0) |
| dataset = processed_dataset[0] |
| else: |
| |
| all_processed_samples = [] |
| for sample in base_dataset: |
| processed_samples = process_dataset(sample) |
| all_processed_samples.extend(processed_samples) |
| |
| processed_dataset = all_processed_samples |
| random.shuffle(processed_dataset) |
| dataset = processed_dataset |
| return dataset |
|
|
|
|
| def get_graph_latent_cot_dataset_backtrack( |
| dataset_path, |
| frontier_stage, |
| r_current, |
| remember_rate, |
| configs, |
| tokenizer, |
| ): |
| """Latent-CoT training set with cross-stage rehearsal / backtracking. |
| |
| Faithful port of the sudoku-latent-backtracking sampler |
| (``_sample_rehearsal_stage``): instead of building every example at the |
| frontier stage, each example is assigned a curriculum stage ``j`` (=latent |
| budget) drawn from a rehearsal distribution: |
| |
| * with prob ``remember_rate``: broad rehearsal, j ~ Uniform{1..frontier} |
| * otherwise: targeted backtracking, j ~ Uniform{r_current..frontier} |
| |
| ``r_current`` is the earliest stage whose retention dropped below the bar |
| (computed from the per-hop eval); when nothing has regressed it equals the |
| frontier, so sampling concentrates on the frontier as in vanilla curriculum. |
| """ |
| base_dataset = json.load(open(dataset_path)) |
|
|
| if configs.debug: |
| base_dataset = base_dataset[:10000] |
| train_size = getattr(configs, "train_size", 0) |
| if train_size and train_size > 0: |
| base_dataset = base_dataset[:train_size] |
|
|
| |
| |
| |
| |
| frontier_stage = max(0, int(frontier_stage)) |
| r_current = min(max(0, int(r_current)), frontier_stage) |
|
|
| def sample_stage(): |
| if random.random() < float(remember_rate): |
| return random.randint(0, frontier_stage) |
| return random.randint(r_current, frontier_stage) |
|
|
| def process_dataset(sample): |
| j = sample_stage() |
| if random.random() < configs.uniform_prob: |
| scheduled_stage_to_train = random.randint(0, min(j, len(sample["steps"]))) |
| else: |
| scheduled_stage_to_train = min(j, len(sample["steps"])) |
|
|
| _smq = bool(getattr(configs, "stage_matched_q", False)) |
| expanded_data = expand_data( |
| sample, |
| scheduled_stage_to_train + 1, |
| len(sample["steps"]), |
| stage_matched_q=_smq, |
| ) |
|
|
| processed_samples = [] |
| for question, continuation in [expanded_data]: |
| question_tokenized = tokenizer.encode(question, add_special_tokens=False) |
| continuation_tokenized = tokenizer.encode(continuation, add_special_tokens=False) |
| tokens = question_tokenized + continuation_tokenized |
| processed_samples.append({ |
| "input_ids": tokens, |
| "labels": [-100] * len(question_tokenized) + continuation_tokenized, |
| "attention_mask": [1] * len(tokens), |
| "position_ids": list(range(len(tokens))), |
| }) |
| return processed_samples |
|
|
| if torch.cuda.device_count() > 1: |
| if dist.get_rank() == 0: |
| all_processed_samples = [] |
| for sample in base_dataset: |
| all_processed_samples.extend(process_dataset(sample)) |
| random.shuffle(all_processed_samples) |
| processed_dataset = [all_processed_samples] |
| else: |
| processed_dataset = [None] |
| dist.broadcast_object_list(processed_dataset, src=0) |
| dataset = processed_dataset[0] |
| else: |
| all_processed_samples = [] |
| for sample in base_dataset: |
| all_processed_samples.extend(process_dataset(sample)) |
| random.shuffle(all_processed_samples) |
| dataset = all_processed_samples |
| return dataset |
|
|
|
|
| def get_graph_latent_question_dataset( |
| dataset_path, |
| scheduled_stage, |
| configs, |
| tokenizer, |
| ): |
| base_dataset = json.load(open(dataset_path)) |
| if configs.debug: |
| base_dataset = base_dataset[:10000] |
| |
| |
| |
| def process_dataset(sample, idx): |
| expanded_data = expand_data(sample, len(sample["steps"]) + 1, len(sample["steps"]), neg_sampling=False) |
| processed_samples = [] |
| for question, continuation in [expanded_data]: |
| question_tokenized = tokenizer.encode(question, add_special_tokens=False) |
| processed_samples.append({ |
| "input_ids": question_tokenized, |
| "attention_mask": [1] * len(question_tokenized), |
| "position_ids": list(range(len(question_tokenized))), |
| "idx": idx, |
| }) |
| return processed_samples |
| |
| if torch.cuda.device_count() > 1: |
| if dist.get_rank() == 0: |
| |
| all_processed_samples = [] |
| for idx, sample in enumerate(base_dataset): |
| processed_samples = process_dataset(sample, idx) |
| all_processed_samples.extend(processed_samples) |
| |
| processed_dataset = all_processed_samples |
| random.shuffle(processed_dataset) |
| processed_dataset = [processed_dataset] |
| else: |
| processed_dataset = [None] |
| dist.broadcast_object_list(processed_dataset, src=0) |
| dataset = processed_dataset[0] |
| else: |
| all_processed_samples = [] |
| for idx, sample in enumerate(base_dataset): |
| all_processed_samples.extend(process_dataset(sample, idx)) |
| random.shuffle(all_processed_samples) |
| dataset = all_processed_samples |
|
|
| return dataset |
|
|
| def get_graph_cot_dataset( |
| dataset_path, |
| configs, |
| tokenizer, |
| ): |
| """ |
| Creates a dataset for training graph reasoning with chain of thought. |
| Each sample will contain the graph edges, question, and the full reasoning path to the answer. |
| """ |
| base_dataset = json.load(open(dataset_path)) |
| |
| if configs.debug: |
| base_dataset = base_dataset[:10000] |
|
|
| def process_dataset(sample): |
| |
| random.shuffle(sample['edges']) |
| |
| |
| question = "<eos> " + "|".join([f" {e[0]} {e[1]} " for e in sample['edges']]).strip() + " [Q] " |
| |
| |
| if random.random() < 0.5: |
| question += f"{sample['target']} {sample['neg_target']}" |
| else: |
| question += f"{sample['neg_target']} {sample['target']}" |
| |
| question += f" [R] {sample['root']}" |
| |
| |
| current_node = sample['root'] |
| continuation = "" |
| for i in range(1, 10): |
| if str(i) in sample["neighbor_k"]: |
| next_node = random.choice( |
| [n for n in sample["neighbor_k"][str(i)] if [current_node, n] in sample['edges']] |
| ) |
| continuation += f" {next_node}" |
| current_node = next_node |
| |
| continuation += f" [A] {sample['target']} <eos>" |
| |
| |
| question_tokenized = tokenizer.encode(question, add_special_tokens=False) |
| continuation_tokenized = tokenizer.encode(continuation, add_special_tokens=False) |
| |
| tokens = question_tokenized + continuation_tokenized |
| |
| processed_sample = { |
| "input_ids": tokens, |
| "labels": [-100] * len(question_tokenized) + continuation_tokenized, |
| "attention_mask": [1] * len(tokens), |
| "position_ids": list(range(len(tokens))), |
| } |
| |
| return [processed_sample] |
|
|
| if torch.cuda.device_count() > 1: |
| if dist.get_rank() == 0: |
| |
| all_processed_samples = [] |
| for sample in base_dataset: |
| processed_samples = process_dataset(sample) |
| all_processed_samples.extend(processed_samples) |
| |
| processed_dataset = all_processed_samples |
| random.shuffle(processed_dataset) |
| processed_dataset = [processed_dataset] |
| else: |
| processed_dataset = [None] |
| dist.broadcast_object_list(processed_dataset, src=0) |
| dataset = processed_dataset[0] |
| else: |
| |
| all_processed_samples = [] |
| for sample in base_dataset: |
| processed_samples = process_dataset(sample) |
| all_processed_samples.extend(processed_samples) |
| |
| processed_dataset = all_processed_samples |
| random.shuffle(processed_dataset) |
| dataset = processed_dataset |
| |
| return dataset |
|
|
| def get_graph_no_cot_dataset( |
| dataset_path, |
| configs, |
| tokenizer, |
| ): |
| """ |
| Creates a dataset for training graph reasoning without chain of thought. |
| Each sample will contain the graph edges, question, and only the final answer |
| without intermediate reasoning steps. |
| """ |
| base_dataset = json.load(open(dataset_path)) |
| |
| if configs.debug: |
| base_dataset = base_dataset[:10000] |
|
|
| def process_dataset(sample): |
| |
| random.shuffle(sample['edges']) |
| |
| |
| question = "<eos> " + "|".join([f" {e[0]} {e[1]} " for e in sample['edges']]).strip() + " [Q] " |
| |
| |
| if random.random() < 0.5: |
| question += f"{sample['target']} {sample['neg_target']}" |
| else: |
| question += f"{sample['neg_target']} {sample['target']}" |
| |
| question += f" [R] {sample['root']}" |
| |
| |
| continuation = f" [A] {sample['target']} <eos>" |
| |
| |
| question_tokenized = tokenizer.encode(question, add_special_tokens=False) |
| continuation_tokenized = tokenizer.encode(continuation, add_special_tokens=False) |
| |
| tokens = question_tokenized + continuation_tokenized |
| |
| processed_sample = { |
| "input_ids": tokens, |
| "labels": [-100] * len(question_tokenized) + continuation_tokenized, |
| "attention_mask": [1] * len(tokens), |
| "position_ids": list(range(len(tokens))), |
| } |
| |
| return [processed_sample] |
|
|
| if torch.cuda.device_count() > 1: |
| if dist.get_rank() == 0: |
| |
| all_processed_samples = [] |
| for sample in base_dataset: |
| processed_samples = process_dataset(sample) |
| all_processed_samples.extend(processed_samples) |
| |
| processed_dataset = all_processed_samples |
| random.shuffle(processed_dataset) |
| processed_dataset = [processed_dataset] |
| else: |
| processed_dataset = [None] |
| dist.broadcast_object_list(processed_dataset, src=0) |
| dataset = processed_dataset[0] |
| else: |
| |
| all_processed_samples = [] |
| for sample in base_dataset: |
| processed_samples = process_dataset(sample) |
| all_processed_samples.extend(processed_samples) |
| |
| processed_dataset = all_processed_samples |
| random.shuffle(processed_dataset) |
| dataset = processed_dataset |
| |
| return dataset |
|
|
| def get_graph_no_latent_question_dataset( |
| dataset_path, |
| configs, |
| tokenizer, |
| ): |
| """ |
| Creates a dataset containing only the questions from the graph reasoning dataset, |
| without any latent tokens. Used for inference to get the input questions. |
| """ |
| base_dataset = json.load(open(dataset_path)) |
| if configs.debug: |
| base_dataset = base_dataset[:10000] |
| |
| def process_dataset(sample, idx): |
| |
| random.shuffle(sample['edges']) |
| question = "<eos> " + "|".join([f" {e[0]} {e[1]} " for e in sample['edges']]).strip() + " [Q] " |
| |
| |
| if random.random() < 0.5: |
| question += f"{sample['target']} {sample['neg_target']}" |
| else: |
| question += f"{sample['neg_target']} {sample['target']}" |
| |
| question += f" [R] {sample['root']}" |
| |
| |
| question_tokenized = tokenizer.encode(question, add_special_tokens=False) |
| |
| processed_sample = { |
| "input_ids": question_tokenized, |
| "attention_mask": [1] * len(question_tokenized), |
| "position_ids": list(range(len(question_tokenized))), |
| "idx": idx, |
| } |
| return [processed_sample] |
| |
| if torch.cuda.device_count() > 1: |
| if dist.get_rank() == 0: |
| |
| all_processed_samples = [] |
| for idx, sample in enumerate(base_dataset): |
| processed_samples = process_dataset(sample, idx) |
| all_processed_samples.extend(processed_samples) |
| |
| processed_dataset = all_processed_samples |
| random.shuffle(processed_dataset) |
| processed_dataset = [processed_dataset] |
| else: |
| processed_dataset = [None] |
| dist.broadcast_object_list(processed_dataset, src=0) |
| dataset = processed_dataset[0] |
| else: |
| |
| all_processed_samples = [] |
| for idx, sample in enumerate(base_dataset): |
| processed_samples = process_dataset(sample, idx) |
| all_processed_samples.extend(processed_samples) |
| |
| processed_dataset = all_processed_samples |
| random.shuffle(processed_dataset) |
| dataset = processed_dataset |
|
|
| return dataset |
|
|
|
|
| def get_graph_finalonly_dataset( |
| dataset_path, |
| scheduled_stage, |
| configs, |
| tokenizer, |
| ): |
| """Final-only variant: at depth d (curriculum), form two depth-d candidates -- one |
| reachable (from neighbor_k), one unreachable (from neg_neighbor_k) -- give d latents, |
| and train the model to output the reachable one as the final answer ([A]). |
| |
| General over any two-component graph: reads only edges/root/neighbor_k/neg_neighbor_k |
| and len(steps). Standard vs BFS flavor is inherited from which frontier the data file |
| stores, exactly like get_graph_latent_cot_dataset.""" |
| base_dataset = json.load(open(dataset_path)) |
|
|
| if configs.debug: |
| base_dataset = base_dataset[:10000] |
|
|
| def process_dataset(sample): |
| L = len(sample["steps"]) |
| d = min(scheduled_stage + 1, L) |
| if random.random() < configs.uniform_prob: |
| d = random.randint(1, d) |
| reach = str(random.choice(sample["neighbor_k"][str(d)])) |
| neg = str(random.choice(sample["neg_neighbor_k"][str(d)])) |
|
|
| edges = list(sample["edges"]) |
| random.shuffle(edges) |
| prefix = "<eos> " + "|".join([f" {e[0]} {e[1]} " for e in edges]).strip() + " [Q] " |
| if random.random() < 0.5: |
| prefix += reach + " " + neg |
| else: |
| prefix += neg + " " + reach |
| prefix += " [R] " + str(sample["root"]) |
| question = prefix + " <|latent|>" * d + " [A] " |
| continuation = reach |
|
|
| question_tokenized = tokenizer.encode(question, add_special_tokens=False) |
| continuation_tokenized = tokenizer.encode(continuation, add_special_tokens=False) |
| tokens = question_tokenized + continuation_tokenized |
| return [{ |
| "input_ids": tokens, |
| "labels": [-100] * len(question_tokenized) + continuation_tokenized, |
| "attention_mask": [1] * len(tokens), |
| "position_ids": list(range(len(tokens))), |
| }] |
|
|
| if torch.cuda.device_count() > 1: |
| if dist.get_rank() == 0: |
| all_processed_samples = [] |
| for sample in base_dataset: |
| all_processed_samples.extend(process_dataset(sample)) |
| processed_dataset = all_processed_samples |
| random.shuffle(processed_dataset) |
| processed_dataset = [processed_dataset] |
| else: |
| processed_dataset = [None] |
| dist.broadcast_object_list(processed_dataset, src=0) |
| dataset = processed_dataset[0] |
| else: |
| all_processed_samples = [] |
| for sample in base_dataset: |
| all_processed_samples.extend(process_dataset(sample)) |
| processed_dataset = all_processed_samples |
| random.shuffle(processed_dataset) |
| dataset = processed_dataset |
| return dataset |
|
|