File size: 5,061 Bytes
f030d3a
 
 
 
 
65db57d
 
f030d3a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
65db57d
f030d3a
65db57d
 
f030d3a
 
 
 
 
 
 
 
 
 
 
fd05733
f030d3a
 
fd05733
f030d3a
 
 
 
 
 
 
 
 
 
 
65db57d
 
 
fd05733
 
f030d3a
 
 
 
 
 
 
 
 
 
 
 
5da4724
 
 
 
 
 
 
 
 
 
 
 
f030d3a
fd05733
 
 
 
 
 
5da4724
 
 
 
 
 
 
 
 
 
fd05733
 
 
 
 
 
f030d3a
65db57d
 
 
 
 
 
 
 
 
 
 
 
 
1d2f1ad
 
65db57d
 
 
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
from textSummarizer.constants import *
from textSummarizer.utils.common import read_yaml, create_directories
from textSummarizer.entity import (DataIngestionConfig,
                                   DataValidationConfig,
                                   DataTransformationConfig,
                                   ModelTrainerConfig,
                                   ModelEvaluationConfig)


class ConfigurationManager:
    def __init__(
        self,
        config_filepath: str = CONFIG_FILE_PATH,
        params_filepath: str = PARAMS_FILE_PATH):
        
        self.config = read_yaml(config_filepath)
        self.params = read_yaml(params_filepath)
        
        create_directories([self.config.artifacts_root])
        
        
    def get_data_ingestion_config(self) -> DataIngestionConfig:
        config =self.config.data_ingestion
        
        create_directories([config.root_dir])
        
        data_ingestion_config = DataIngestionConfig(
            root_dir = Path(config.root_dir),
            source_URL = config.source_URL,
            local_data_file = Path(config.local_data_file),
            unzip_dir = Path(config.unzip_dir),
        )
        
        return data_ingestion_config
    
    
    def get_data_validation_config(self) -> DataValidationConfig:
        config = self.config.data_validation
        
        create_directories([config.root_dir])
        
        data_validation_config = DataValidationConfig(
            root_dir=Path(config.root_dir),
            STATUS_FILE=config.STATUS_FILE,
            ALL_REQUIRED_FILES=config.ALL_REQUIRED_FILES,
            data_dir=Path(config.data_dir),
        )
        
        return data_validation_config
    
    
    def get_data_transformation_config(self) -> DataTransformationConfig:
        config= self.config.data_transformation
        
        create_directories([config.root_dir])
        
        data_transformation_config = DataTransformationConfig(
            root_dir=Path(config.root_dir),
            data_path=Path(config.data_path),
            tokenizer_name=config.tokenizer_name,  # if this is a path, use Path(); if just a model name, keep as str
            dev_run=getattr(config, 'dev_run', False),
            dev_model=getattr(config, 'dev_model', None),
        )
        
        return data_transformation_config
    
    
    def get_model_trainer_config(self) -> ModelTrainerConfig:
        config = self.config.model_trainer
        params = self.params.TrainingArguments
        
        create_directories([config.root_dir])
        
        model_trainer_config = ModelTrainerConfig(
            root_dir=Path(config.root_dir),
            data_path=Path(config.data_path),
            model_ckpt=str(config.model_ckpt),
            num_train_epochs=int(params.num_train_epochs),
            warmup_steps=int(params.warmup_steps),
            per_device_train_batch_size=int(params.per_device_train_batch_size),
            weight_decay=float(params.weight_decay),
            logging_steps=int(params.logging_steps),
            eval_strategy=str(params.eval_strategy),
            eval_steps=int(params.eval_steps),
            save_steps=int(float(params.save_steps)),
            gradient_accumulation_steps=int(params.gradient_accumulation_steps)
        )

        # Optional dev quick-train settings (not all configs will have these)
        try:
            model_trainer_config = ModelTrainerConfig(
                root_dir=Path(config.root_dir),
                data_path=Path(config.data_path),
                model_ckpt=str(config.model_ckpt),
                num_train_epochs=int(params.num_train_epochs),
                warmup_steps=int(params.warmup_steps),
                per_device_train_batch_size=int(params.per_device_train_batch_size),
                weight_decay=float(params.weight_decay),
                logging_steps=int(params.logging_steps),
                eval_strategy=str(params.eval_strategy),
                eval_steps=int(params.eval_steps),
                save_steps=int(float(params.save_steps)),
                gradient_accumulation_steps=int(params.gradient_accumulation_steps),
                dev_run=getattr(config, 'dev_run', False),
                dev_model=getattr(config, 'dev_model', None),
                dev_subset=int(getattr(config, 'dev_subset', 0))
            )
        except Exception:
            pass
        
        return model_trainer_config
    
    
    def get_model_evaluation_config(self) -> ModelEvaluationConfig:
        config = self.config.model_evaluation

        create_directories([config.root_dir])

        model_evaluation_config = ModelEvaluationConfig(
            root_dir=config.root_dir,
            data_path=config.data_path,
            model_path = config.model_path,
            tokenizer_path = config.tokenizer_path,
            metric_file_name = config.metric_file_name,
            hub_model_id = getattr(config, "hub_model_id", "google/pegasus-cnn_dailymail"),
        )

        return model_evaluation_config