File size: 2,670 Bytes
f030d3a
 
 
 
 
 
 
 
 
fd05733
 
 
 
 
 
f030d3a
 
65db57d
 
1d2f1ad
65db57d
 
 
 
 
 
 
f030d3a
 
 
 
 
 
 
23e50b1
 
 
 
 
 
 
 
65db57d
 
 
f030d3a
fd05733
 
 
 
 
 
 
65db57d
 
23e50b1
 
 
 
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
import os
from textSummarizer.logging import logger
from transformers import AutoTokenizer
from datasets import load_dataset, load_from_disk
from textSummarizer.entity import DataTransformationConfig

class DataTransformation:
    def __init__(self, config: DataTransformationConfig):
        self.config = config
        # choose tokenizer checkpoint: prefer dev_model when dev_run is enabled
        tokenizer_checkpoint = self.config.tokenizer_name
        if getattr(self.config, 'dev_run', False) and getattr(self.config, 'dev_model', None):
            tokenizer_checkpoint = self.config.dev_model

        self.tokenizer = AutoTokenizer.from_pretrained(tokenizer_checkpoint)
        
    def convert_examples_to_features(self, example_batch):
        input_encodings = self.tokenizer(
            example_batch['dialogue'],
            max_length=512,
            truncation=True
        )
        target_encodings = self.tokenizer(
            text_target=example_batch['summary'],
            max_length=128,
            truncation=True
        )
        return {
            'input_ids': input_encodings['input_ids'],
            'attention_mask': input_encodings['attention_mask'],
            'labels': target_encodings['input_ids'],
        }
        
    def convert(self):
        save_path = self.config.root_dir / "samsum_dataset"
        
        # Temporarily disabled skip check to force transformation
        # # Check if processed dataset already exists
        # if save_path.exists() and (save_path / "dataset_dict.json").exists():
        #     logger.info(f"Processed dataset already exists at {save_path}. Skipping transformation.")
        #     return
            
        logger.info(f"Loading dataset from {self.config.data_path}")
        dataset_samsum = load_from_disk(str(self.config.data_path))
        logger.info("Tokenizing dataset...")
        dataset_samsum_pt = dataset_samsum.map(self.convert_examples_to_features, batched=True)
        
        # ensure labels are correctly formatted for seq2seq Trainer
        def rename_for_trainer(batch):
            # already returns 'labels' in convert_examples_to_features; keep
            return batch

        dataset_samsum_pt = dataset_samsum_pt.map(rename_for_trainer, batched=True)
        os.makedirs(save_path, exist_ok=True)
        logger.info(f"Saving processed dataset to {save_path}")
        # Use absolute path with proper Windows formatting to handle spaces in path
        abs_save_path = os.path.abspath(str(save_path))
        # Disable multiprocessing on Windows to avoid path issues with spaces
        dataset_samsum_pt.save_to_disk(abs_save_path, num_proc=1)