File size: 32,530 Bytes
1ac68b2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
632
633
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
655
656
657
658
659
660
661
662
663
664
665
666
667
668
669
670
671
672
673
674
675
676
677
678
679
680
681
682
683
684
685
686
687
688
689
690
691
692
693
694
695
696
697
698
699
700
701
702
703
704
705
706
707
708
709
710
711
712
713
714
715
716
717
718
719
720
721
722
723
724
725
726
727
728
729
730
731
732
733
734
735
736
737
738
739
740
741
742
743
744
745
746
747
748
749
750
751
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