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() @abstractmethod 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): @override 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