AI_Summarizer / src /textSummarizer /config /configuration.py
Jeevant10's picture
deploy
1d2f1ad
Raw
History Blame Contribute Delete
5.06 kB
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