Download data_module.py from Android12138/FlowResampler: direct link, hf CLI and curl.
- Browser
- Download file 32.5 kB
-
https://huggingface.co/Android12138/FlowResampler/resolve/main/data_module.py
- Command line
-
hf download hf://Android12138/FlowResampler/data_module.py
-
curl -L -o data_module.py https://huggingface.co/Android12138/FlowResampler/resolve/main/data_module.py
32.5 kB
| import os | |
| import copy | |
| from tqdm import tqdm | |
| from abc import ABC, abstractmethod | |
| from typing import override, Union | |
| import torch | |
| from torch.utils.data import Dataset, DataLoader | |
| from torch_geometric.data import Data, Batch | |
| from sentence_transformers import SentenceTransformer | |
| from transformers import AutoTokenizer, DataCollatorWithPadding | |
| from trl.trainer.sft_trainer import DataCollatorForLanguageModeling | |
| from utils.json_utils import load_json, load_jsonl | |
| class GraphChatDataset(Dataset): | |
| """A custom PyTorch Dataset to handle mixed PyTorch Geometric Data and text prompts. | |
| It avoids PyArrow serialization issues by keeping everything in native Python lists. | |
| """ | |
| def __init__(self, data_list: list[dict]): | |
| self.data_list = data_list | |
| def __len__(self): | |
| return len(self.data_list) | |
| def __getitem__(self, idx): | |
| # Return a shallow copy of the dictionary. | |
| # This is CRITICAL because `collate_fn` uses `.pop()`, | |
| # which would otherwise mutate the underlying data and cause errors in the 2nd epoch. | |
| return {k: v for k, v in self.data_list[idx].items()} | |
| def map(self, fn, fn_kwargs=None, batched=False, remove_columns=None): | |
| """Mimics the Hugging Face dataset.map() interface for seamless integration.""" | |
| if fn_kwargs is None: | |
| fn_kwargs = {} | |
| new_data_list = [] | |
| for item in tqdm(self.data_list, desc="Mapping dataset"): | |
| # Apply the tokenize_fn | |
| new_item = fn(item, **fn_kwargs) | |
| # Since tokenize_fn explicitly returns the exact dictionary structure needed, | |
| # we can safely ignore `remove_columns` and just append the returned item. | |
| new_data_list.append(new_item) | |
| return GraphChatDataset(new_data_list) | |
| # ========================================== | |
| # Tokenization Strategies | |
| # ========================================== | |
| class ClsTokenizeStrategy: | |
| def __init__(self, tokenizer: AutoTokenizer = None): | |
| self.tokenizer = tokenizer | |
| def __call__(self, example): | |
| """Tokenizes text into input_ids and attention_mask.""" | |
| return self.tokenizer( | |
| example["funcs"], | |
| truncation=True, | |
| padding=False, | |
| max_length=None, | |
| ) | |
| class CompletionTokenizeStrategy: | |
| def __init__(self, tokenizer: AutoTokenizer, max_length=None, data_formats: str="pyg_graph"): | |
| self.tokenizer = tokenizer | |
| self.max_length = max_length | |
| # Route tokenization logic based on data_formats during initialization | |
| if data_formats == "pyg_graph": | |
| self._tokenize_fn = self._tokenize_pyg_graph | |
| else: | |
| self._tokenize_fn = self._tokenize_text_only | |
| def __call__(self, example): | |
| return self._tokenize_fn(example) | |
| def _base_tokenize(self, example): | |
| """Tokenizes a single example into input_ids, attention_mask, and labels.""" | |
| prompt = example["prompts"] | |
| completion = example["completions"] | |
| # 1. Construct the "question only" message list -> to calculate prompt length | |
| prompt_ids = self.tokenizer.apply_chat_template( | |
| prompt, | |
| tokenize=True, | |
| add_generation_prompt=True, | |
| truncation=False, | |
| enable_thinking=False, | |
| ).input_ids | |
| # 2. Construct the "full conversation" message list -> input_ids | |
| input_ids = self.tokenizer.apply_chat_template( | |
| prompt+completion, | |
| tokenize=True, | |
| truncation=False, | |
| enable_thinking=False, | |
| ).input_ids | |
| # 3. Generate labels and apply masking | |
| labels = copy.deepcopy(input_ids) | |
| prompt_len = len(prompt_ids) | |
| # Set the labels of the prompt tokens to -100 to ignore them in loss computation | |
| for i in range(len(labels)): | |
| if i < prompt_len: | |
| labels[i] = -100 | |
| # 4. Truncate if max_length is specified | |
| if not self.max_length is None: | |
| if len(input_ids) > self.max_length: | |
| input_ids = input_ids[:self.max_length] | |
| labels = labels[:self.max_length] # no need to shift, when feed to hugging face CausalLM, it will automatically shift the labels by one to the left internally | |
| # 5. Create attention mask and position ids | |
| attention_mask = [1] * len(input_ids) | |
| return { | |
| "input_ids": input_ids, | |
| "attention_mask": attention_mask, | |
| "labels": labels, | |
| "indices": example["indices"], | |
| } | |
| def _tokenize_pyg_graph(self, example): | |
| features = self._base_tokenize(example) | |
| features["graphs"] = example["graphs"] | |
| return features | |
| def _tokenize_text_only(self, example): | |
| return self._base_tokenize(example) | |
| class PromptTokenizeStrategy: | |
| def __init__(self, tokenizer: AutoTokenizer, answer_template: str, max_length=None, data_formats: str = "pyg_graph"): | |
| self.tokenizer = tokenizer | |
| self.answer_template = answer_template | |
| self.max_length = max_length | |
| # Route tokenization logic based on data_formats during initialization | |
| if data_formats == "pyg_graph": | |
| self._tokenize_fn = self._tokenize_pyg_graph | |
| else: | |
| self._tokenize_fn = self._tokenize_text_only | |
| def __call__(self, example): | |
| return self._tokenize_fn(example) | |
| def _base_tokenize(self, example): | |
| """Tokenizes a single example into prompt_ids and prompt_attention_mask.""" | |
| prompt = example["prompts"] | |
| # Construct the "question only" message list -> to calculate prompt length | |
| prompt_ids = self.tokenizer.apply_chat_template( | |
| prompt, | |
| tokenize=True, | |
| add_generation_prompt=True, | |
| truncation=False, | |
| enable_thinking=False, | |
| ).input_ids | |
| # Create prompt_ids for constrained decoding | |
| answer_ids = self.tokenizer(self.answer_template, add_special_tokens=False).input_ids | |
| prompt_ids = prompt_ids + answer_ids | |
| prompt_attention_mask = [1] * len(prompt_ids) | |
| return { | |
| "prompts": prompt_ids, | |
| "prompt_attention_mask": prompt_attention_mask, | |
| "indices": example["indices"], | |
| } | |
| def _tokenize_pyg_graph(self, example): | |
| features = self._base_tokenize(example) | |
| features["graphs"] = example["graphs"] | |
| return features | |
| def _tokenize_text_only(self, example): | |
| return self._base_tokenize(example) | |
| # ========================================== | |
| # Collation Strategies | |
| # ========================================== | |
| class ClsDataCollator: | |
| def __init__(self, tokenizer: AutoTokenizer = None, data_formats: str = "pyg_graph"): | |
| # Route collation logic based on data_formats during initialization | |
| if data_formats == "no_graph": | |
| self.data_collator = DataCollatorWithPadding( | |
| tokenizer=tokenizer, | |
| padding="longest", | |
| ) | |
| self._collate_fn = self._collate_text_only | |
| else: | |
| self._collate_fn = self._collate_pyg_graph | |
| def __call__(self, batch): | |
| return self._collate_fn(batch) | |
| def _collate_text_only(self, batch): | |
| """Custom collate function to handle text and indices""" | |
| indices = [item.pop("indices") for item in batch] | |
| labels = [item.pop("labels") for item in batch] | |
| # Collate input_ids and attention_mask with dynamic padding | |
| batch = self.data_collator(batch) | |
| # Put the indices and labels back into the batch | |
| batch["indices"] = indices | |
| batch["labels"] = torch.tensor(labels, dtype=torch.long) | |
| return batch | |
| def _collate_pyg_graph(self, batch): | |
| """Custom collate function to handle graphs and indices""" | |
| # Extract fields to process them separately | |
| indices = [item.pop("indices") for item in batch] | |
| graphs = [item.pop("graphs") for item in batch] | |
| labels = [item.pop("labels") for item in batch] | |
| # Reassemble the batch | |
| batch_dict = { | |
| "indices": indices, | |
| "labels": torch.tensor(labels, dtype=torch.long), | |
| "graphs": Batch.from_data_list(graphs) | |
| } | |
| return batch_dict | |
| class CompletionDataCollator: | |
| def __init__(self, tokenizer: AutoTokenizer, data_formats: str = "pyg_graph"): | |
| # DataCollatorForLanguageModeling uses right-padding only. | |
| # During SFT, both `input_ids` and `labels` are fed directly to `model.forward()` to compute the loss. | |
| # Since `model.forward()` does not automatically skip padding tokens when generating `position_ids`, | |
| # it simply assigns them sequentially based on the entire sequence length. | |
| # Therefore, right-padding is required to keep the `position_ids` of the valid text aligned across the batch. | |
| self.data_collator = DataCollatorForLanguageModeling( | |
| pad_token_id=tokenizer.pad_token_id, | |
| padding_free=False, | |
| pad_to_multiple_of=None, | |
| return_tensors="pt", | |
| ) | |
| # Route collation logic based on data_formats during initialization | |
| if data_formats == "pyg_graph": | |
| self._collate_fn = self._collate_pyg_graph | |
| else: | |
| self._collate_fn = self._collate_text_only | |
| def __call__(self, example): | |
| return self._collate_fn(example) | |
| def _collate_pyg_graph(self, batch): | |
| """Custom collate function to handle mixed data types (indices, graphs, and texts)""" | |
| # Extract non-standard fields to process them separately | |
| indices = [item.pop("indices") for item in batch] | |
| graphs = [item.pop("graphs") for item in batch] | |
| # Process input_ids, attention_mask, and labels with DataCollatorForLanguageModeling as labels needs to be padded with -100 | |
| completions = self.data_collator(batch) | |
| # Reassemble the batch | |
| completions["indices"] = indices | |
| completions["graphs"] = Batch.from_data_list(graphs) | |
| return completions | |
| def _collate_text_only(self, batch): | |
| """Custom collate function to handle mixed data types (indices and texts)""" | |
| # Extract non-standard fields to process them separately | |
| indices = [item.pop("indices") for item in batch] | |
| # Process input_ids, attention_mask, and labels with DataCollatorForLanguageModeling as labels needs to be padded with -100 | |
| completions = self.data_collator(batch) | |
| # Reassemble the batch | |
| completions["indices"] = indices | |
| return completions | |
| class PromptDataCollator: | |
| def __init__(self, tokenizer: AutoTokenizer, data_formats: str = "pyg_graph"): | |
| # DataCollatorWithPadding pads the sequences based on the tokenizer settings. | |
| # --------------------------------------------------- | |
| # In `model.generate()`, padding tokens are automatically ignored when computing `position_ids`. | |
| # This prevents `position_ids` from being misaligned by leading pad tokens during generation. | |
| # Left-padding is strictly required for batched generation to ensure the causal LM always | |
| # predicts the next token from the last valid (non-pad) token. | |
| # --------------------------------------------------- | |
| # If `model.forward()` is explicitly used to fetch next-token logits instead of `generate()`, | |
| # left-padding MUST still be used. With right-padding, `logits[:, -1, :]` would incorrectly | |
| # point to a padding token rather than the actual end of the text sequence. | |
| # --------------------------------------------------- | |
| # In our case, because we need to concat prefix and get next-token logits during inference. | |
| # For simplicity, we use a batch size of 1 during testing, though batching is also supported. | |
| self.data_collator = DataCollatorWithPadding( | |
| tokenizer=tokenizer, | |
| padding="longest" | |
| ) | |
| # Route collation logic based on data_formats during initialization | |
| if data_formats == "pyg_graph": | |
| self._collate_fn = self._collate_pyg_graph | |
| else: | |
| self._collate_fn = self._collate_text_only | |
| def __call__(self, example): | |
| return self._collate_fn(example) | |
| def _collate_pyg_graph(self, batch): | |
| """Custom collate function to handle mixed data types (indices, graphs, and texts)""" | |
| # Extract non-standard fields to process them separately | |
| indices = [item.pop("indices") for item in batch] | |
| graphs = [item.pop("graphs") for item in batch] | |
| # Process prompts and prompt_attention_mask with DataCollatorWithPadding | |
| prompts = self.data_collator([{"input_ids": item.pop("prompts"), "attention_mask": item.pop("prompt_attention_mask")} for item in batch]) | |
| # Reassemble the batch | |
| return { | |
| "prompts": prompts["input_ids"], | |
| "prompt_attention_mask": prompts["attention_mask"], | |
| "indices": indices, | |
| "graphs": Batch.from_data_list(graphs), | |
| } | |
| def _collate_text_only(self, batch): | |
| """Custom collate function to handle mixed data types (indices and texts)""" | |
| # Extract non-standard fields to process them separately | |
| indices = [item.pop("indices") for item in batch] | |
| # Process prompts and prompt_attention_mask with DataCollatorWithPadding | |
| prompts = self.data_collator([{"input_ids": item.pop("prompts"), "attention_mask": item.pop("prompt_attention_mask")} for item in batch]) | |
| # Reassemble the batch | |
| return { | |
| "prompts": prompts["input_ids"], | |
| "prompt_attention_mask": prompts["attention_mask"], | |
| "indices": indices, | |
| } | |
| # ========================================== | |
| # Data Module | |
| # ========================================== | |
| class DataModule(ABC): | |
| def __init__( | |
| self, | |
| config, | |
| ): | |
| """Initialize the DataModule for dataset processing. | |
| Args: | |
| config (BaseConfig): Configuration object containing parameters for data processing. | |
| """ | |
| # Initialize the config | |
| self.config = config | |
| # Initialize the tokenizer to filter samples by max tokens and convert the graph text into tokens | |
| self.tokenizer = AutoTokenizer.from_pretrained( | |
| config.func_filter_model_name, | |
| padding_side="left", | |
| ) | |
| # 1. initialize the dataset and directories | |
| self._prepare_dataset_and_directory() | |
| # 2. cut into subsets | |
| self._generate_splits() | |
| def _prepare_dataset_and_directory(self): | |
| """TODO: Prepare the raw dataset and create save directories.""" | |
| pass | |
| def _generate_splits(self): | |
| """Generates train, validation, and test splits based on the valid_func_length_indices. | |
| If the splits already exist, it loads them from the file. | |
| If not, it generates the splits and saves them to the file. | |
| """ | |
| print("Loading train, val, test indices from file:", f"{self.indices_directory}/test_indices.json") | |
| self.train_indices = load_json(f"{self.indices_directory}/train_indices.json") | |
| self.val_indices = load_json(f"{self.indices_directory}/val_indices.json") | |
| self.test_indices = load_json(f"{self.indices_directory}/test_indices.json") | |
| def graph_to_context(self, graph: dict) -> tuple[str, str]: | |
| """ | |
| Convert node and edge information from a graph as context. | |
| Args: | |
| graph (Data): A PyTorch Geometric Data object representing the graph. | |
| Returns: | |
| tuple[str, str]: Node and edge information in string format. | |
| """ | |
| node_info = "Node ID\t Node Type\t CODE\n" | |
| for i, (label, code) in enumerate(zip(graph.node_label, graph.CODE)): | |
| node_info += f"{i}\t {label}\t {code}\n" | |
| edge_info = "Source\t Target\t Edge Type\n" | |
| for i in range(graph.edge_index.size(1)): | |
| source = graph.edge_index[0, i].item() | |
| target = graph.edge_index[1, i].item() | |
| label = graph.edge_label[i] | |
| edge_info += f"{source}\t {target}\t {label}\n" | |
| return node_info, edge_info | |
| class PrimeVulDataModule(DataModule): | |
| def _prepare_dataset_and_directory(self): | |
| """Prepares the PrimeVul dataset by loading cwe info and formatting into paired dicts. | |
| Each pair consists of two consecutive dicts, where the first is vulnerable and second is safe. | |
| """ | |
| # Initialize the directories | |
| self.indices_directory = f"datasets/processed/primevul/data_splits/" | |
| # Load the dataset and graph | |
| self.dataset = ( | |
| load_jsonl("datasets/raw/primevul/primevul_train.jsonl") + | |
| load_jsonl("datasets/raw/primevul/primevul_valid.jsonl") + | |
| load_jsonl("datasets/raw/primevul/primevul_test.jsonl") | |
| ) | |
| # Conditionally load and process graphs | |
| graphs = self._prepare_pyg_graphs() | |
| # Add graph and index into the dataset | |
| for i, item in enumerate(self.dataset): | |
| item["index"] = i | |
| if graphs is not None: | |
| item["graph"] = graphs[i] | |
| # Initialize the index to label mapping | |
| self.index_to_label = {1: "true", 0: "false"} # 1 is vulnerable, 0 is secure | |
| def _prepare_pyg_graphs(self) -> Union[list[Data], None]: | |
| """Prepares PyTorch Geometric Data objects from the raw graphs in JSONL format. | |
| """ | |
| if self.config.disable_graph_data: | |
| return None | |
| print("Loading raw graphs from JSONL files...") | |
| graphs = ( | |
| load_jsonl("datasets/processed/primevul/data_graphs/cpg/processed/primevul_train.jsonl") + | |
| load_jsonl("datasets/processed/primevul/data_graphs/cpg/processed/primevul_valid.jsonl") + | |
| load_jsonl("datasets/processed/primevul/data_graphs/cpg/processed/primevul_test.jsonl") | |
| ) | |
| # If the embedding model is configured, generate node and edge embeddings | |
| embedding_enabled = getattr(self.config, "gnn_embedding_model_name", None) is not None | |
| if embedding_enabled: | |
| # Create embedding save directory if it doesn't exist | |
| embedding_dir = "datasets/processed/primevul/data_graphs/cpg/embeddings" | |
| os.makedirs(embedding_dir, exist_ok=True) | |
| node_emb_path = os.path.join(embedding_dir, "node_embeddings.pt") | |
| edge_map_path = os.path.join(embedding_dir, "edge_label_to_attr.pt") | |
| # Check if the saved embeddings and mapping already exist | |
| if os.path.exists(node_emb_path) and os.path.exists(edge_map_path): | |
| print(f"Loading cached node embeddings from {node_emb_path}...") | |
| node_embeddings = torch.load(node_emb_path, map_location="cpu", weights_only=True) | |
| print(f"Loading cached edge_label_to_attr mapping from {edge_map_path}...") | |
| edge_label_to_attr = torch.load(edge_map_path, map_location="cpu", weights_only=True) | |
| else: | |
| print("Cached embeddings not found. Initializing the embedding model...") | |
| embedding_model = SentenceTransformer( | |
| self.config.gnn_embedding_model_name, | |
| trust_remote_code=True, | |
| model_kwargs={"torch_dtype": self.config.gnn_dtype}, | |
| ) | |
| # Generate global edge mapping based on all graphs | |
| print("Extracting unique edge labels across all graphs...") | |
| unique_edge_labels = set() | |
| for graph in graphs: | |
| # Extract labels from the 'edges' list | |
| unique_edge_labels.update([edge["label"] for edge in graph["edges"]]) | |
| unique_edge_labels = list(unique_edge_labels) | |
| edge_label_to_attr = {} | |
| if unique_edge_labels: | |
| print("Generating embeddings for unique edge labels...") | |
| edge_embs = embedding_model.encode( | |
| unique_edge_labels, | |
| truncate_dim=self.config.gnn_edge_dim, | |
| convert_to_tensor=True, | |
| show_progress_bar=False, | |
| ) | |
| # Map string label to CPU tensor | |
| edge_label_to_attr = { | |
| label: emb.cpu() for label, emb in zip(unique_edge_labels, edge_embs) | |
| } | |
| # Generate node embeddings for all graphs | |
| node_embeddings = [] | |
| for graph in tqdm(graphs, desc="Generating node embeddings"): | |
| # Extract CODE and label from the 'nodes' list | |
| node_inputs = [f"{node['CODE']} {node['label']}".strip() for node in graph["nodes"]] | |
| # Generate node embeddings | |
| x = embedding_model.encode( | |
| node_inputs, | |
| truncate_dim=self.config.gnn_input_size, | |
| convert_to_tensor=True, | |
| show_progress_bar=False, | |
| ) | |
| node_embeddings.append(x.cpu()) | |
| # Save the generated node embeddings and edge mapping | |
| print("Saving generated node embeddings and edge mapping to disk...") | |
| torch.save(node_embeddings, node_emb_path) | |
| torch.save(edge_label_to_attr, edge_map_path) | |
| # Inject embeddings and convert directly to PyG Data objects | |
| pyg_graphs = [] | |
| for i, graph in enumerate(tqdm(graphs, desc="Converting to PyG Data")): | |
| # Map original node IDs to continuous indices (0, 1, 2...) as required by PyG | |
| id_to_idx = {node["id"]: idx for idx, node in enumerate(graph["nodes"])} | |
| # Construct the PyG edge_index tensor with shape [2, num_edges] | |
| src = [id_to_idx[edge["source"]] for edge in graph["edges"]] | |
| dst = [id_to_idx[edge["target"]] for edge in graph["edges"]] | |
| edge_index = torch.tensor([src, dst], dtype=torch.long) | |
| # Basic PyG attributes | |
| data_kwargs = { | |
| "edge_index": edge_index, | |
| "CODE": [node["CODE"] for node in graph["nodes"]], | |
| "node_label": [node["label"] for node in graph["nodes"]], | |
| "edge_label": [edge["label"] for edge in graph["edges"]], | |
| } | |
| # Add embeddings only when enabled | |
| if embedding_enabled: | |
| data_kwargs["x"] = node_embeddings[i][:, :self.config.gnn_input_size] | |
| edge_attr = torch.stack([edge_label_to_attr[edge["label"]] for edge in graph["edges"]]) | |
| data_kwargs["edge_attr"] = edge_attr[:, :self.config.gnn_edge_dim] | |
| # Construct the PyG Data object | |
| pyg_graphs.append(Data(**data_kwargs)) | |
| return pyg_graphs | |
| def get_cls_dataset( | |
| self, | |
| indices: list, | |
| data_formats: str = "no_graph", | |
| ): | |
| """Returns a dataset based on the provided indices and data type. | |
| Args: | |
| indices (list): List of indices to select from the dataset. | |
| data_formats (str): Data formulation mode. Options: | |
| - "no_graph": Uses only func (no graph data). | |
| - "pyg_graph": Returns PyG graph objects. | |
| - "text_graph": Returns graph context (text) and func. | |
| Returns: | |
| dataset (List): A list containing the selected items and their attributes. | |
| """ | |
| # Define extraction functions to handle different data extraction logic | |
| if data_formats == "no_graph": | |
| def extract_fn(item): | |
| return {"funcs": item["func"]} | |
| elif data_formats == "pyg_graph": | |
| def extract_fn(item): | |
| return {"graphs": item["graph"]} | |
| elif data_formats == "text_graph": | |
| def extract_fn(item): | |
| node_info, edge_info = self.graph_to_context(item["graph"]) | |
| return { | |
| "node_info": node_info, | |
| "edge_info": edge_info, | |
| "funcs": item["func"], | |
| } | |
| dataset = [] | |
| for index in tqdm(indices): | |
| item = self.dataset[index] | |
| # Assemble common fields | |
| data_dict = { | |
| "indices": item["index"], | |
| "labels": item["target"], | |
| } | |
| # Assemble specific fields | |
| data_dict.update(extract_fn(item)) | |
| dataset.append(data_dict) | |
| return dataset | |
| def get_chat_dataset( | |
| self, | |
| indices: list, | |
| question_template: str, | |
| answer_template: str, | |
| data_formats: str = "pyg_graph", | |
| ): | |
| """Prepares a Dataset object in the format of prompts and completions for the given indices. | |
| Args: | |
| indices (list): List of indices to select from the dataset. | |
| question_template (str): The template for the question prompt. | |
| answer_template (str): The template for the answer prompt. | |
| data_formats (str): Data formulation mode. Options: | |
| - "no_graph": Uses only func in the prompt (no graph data). | |
| - "pyg_graph": Returns PyG graph objects + func in the prompt. | |
| - "text_graph": Uses graph context (text) + func in the prompt. | |
| Returns: | |
| dataset (Dataset): A list containing the selected items and their attributes. | |
| """ | |
| # Define extraction functions to handle different data extraction logic | |
| if data_formats == "text_graph": | |
| def extract_fn(item): | |
| return { | |
| "prompts": [{ | |
| "content": question_template.format( | |
| func=item["func"], | |
| graph=self.graph_to_context(item["graph"]), | |
| ), | |
| "role": "user", | |
| }] | |
| } | |
| elif data_formats == "pyg_graph": | |
| def extract_fn(item): | |
| return { | |
| "graphs": item["graph"], | |
| "prompts": [{ | |
| "content": question_template.format( | |
| func=item["func"], | |
| ), | |
| "role": "user", | |
| }] | |
| } | |
| elif data_formats == "no_graph": | |
| def extract_fn(item): | |
| return { | |
| "prompts": [{ | |
| "content": question_template.format( | |
| func=item["func"], | |
| ), | |
| "role": "user", | |
| }] | |
| } | |
| dataset = [] | |
| for index in tqdm(indices): | |
| item = self.dataset[index] | |
| # Assemble common fields | |
| data_dict = { | |
| "indices": item["index"], # used for result logging and loading, this needs to be a string. | |
| "completions": [ | |
| { | |
| "content": f"{answer_template}{self.index_to_label[item['target']]}", | |
| "role": "assistant", | |
| }, | |
| ] | |
| } | |
| # Execute the externally bound function | |
| data_dict.update(extract_fn(item)) | |
| dataset.append(data_dict) | |
| return dataset | |
| def get_dataloader( | |
| self, | |
| task_types: str, | |
| indices: list, | |
| batch_size: int, | |
| tokenizer: AutoTokenizer = None, | |
| data_formats: str = "pyg_graph", | |
| question_template: str = None, | |
| answer_template: str = None, | |
| ): | |
| """Returns a unified DataLoader based on the provided task type and indices. | |
| Args: | |
| task_types (str): The specific task formulation. Options: | |
| - "cls": Classification task format. | |
| - "completion": Completion-only format for chat models. | |
| - "prompt": Prompt-only format for chat models. | |
| indices (List): List of indices to select from the dataset. | |
| batch_size (int): The batch size for the DataLoader. | |
| tokenizer (AutoTokenizer, optional): The tokenizer to use. Required for text processing. | |
| data_formats (str): Data formulation mode. Options: | |
| - "no_graph": Uses text data only. | |
| - "pyg_graph": Uses PyTorch Geometric graph objects. | |
| - "text_graph": Uses text-based graph context (applicable for chat tasks). | |
| question_template (str, optional): The template for the question prompt (used in chat tasks). | |
| answer_template (str, optional): The template for the answer prompt (used in chat tasks). | |
| Returns: | |
| dataloader (DataLoader): A unified PyTorch DataLoader object configured for the specified task. | |
| """ | |
| # Define the strategy mapping for different task types | |
| strategy_map = { | |
| "cls": { | |
| "dataset_fn": lambda: self.get_cls_dataset(indices, data_formats), | |
| "tokenizer_cls": ClsTokenizeStrategy if data_formats == "no_graph" else None, | |
| "collator_cls": ClsDataCollator, | |
| "remove_cols": ["funcs"] if data_formats == "no_graph" else [], | |
| }, | |
| "completion": { | |
| "dataset_fn": lambda: self.get_chat_dataset(indices, question_template, answer_template, data_formats), | |
| "tokenizer_cls": CompletionTokenizeStrategy, | |
| "collator_cls": CompletionDataCollator, | |
| "remove_cols": ["prompts", "completions"], | |
| }, | |
| "prompt": { | |
| "dataset_fn": lambda: self.get_chat_dataset(indices, question_template, answer_template, data_formats), | |
| "tokenizer_cls": PromptTokenizeStrategy, | |
| "collator_cls": PromptDataCollator, | |
| "remove_cols": ["prompts", "completions"], | |
| } | |
| } | |
| if task_types not in strategy_map: | |
| raise ValueError(f"Invalid task_types: '{task_types}'. Choose from {list(strategy_map.keys())}") | |
| cfg = strategy_map[task_types] | |
| # Retrieve and wrap the dataset | |
| dataset = GraphChatDataset(cfg["dataset_fn"]()) | |
| # Apply tokenization conditionally (skips mapping for PyG graphs in cls tasks) | |
| if cfg["tokenizer_cls"] is not None: | |
| # Prepare arguments specifically required by the tokenization strategies | |
| tokenize_kwargs = {"tokenizer": tokenizer, "data_formats": data_formats} | |
| if task_types in ["completion", "prompt"]: | |
| tokenize_kwargs["max_length"] = None | |
| if task_types == "prompt": | |
| tokenize_kwargs["answer_template"] = answer_template | |
| # Instantiate the strategy and map it to the dataset | |
| tokenize_fn = cfg["tokenizer_cls"](**tokenize_kwargs) | |
| dataset = dataset.map( | |
| tokenize_fn, | |
| batched=False, # Process sample by sample due to complex tokenization logic | |
| remove_columns=cfg["remove_cols"], | |
| ) | |
| # Instantiate the data collator and generate the DataLoader | |
| collate_fn = cfg["collator_cls"]( | |
| tokenizer=tokenizer, | |
| data_formats=data_formats, | |
| ) | |
| dataloader = DataLoader( | |
| dataset, | |
| batch_size=batch_size, | |
| shuffle=False, # no shuffle as train_test_split already handles it | |
| collate_fn=collate_fn, | |
| pin_memory=True, | |
| ) | |
| return dataloader | |