FirenetCNN / src /training.py
OpelSpeedster's picture
Update the Project
1feed70 verified
Raw
History Blame Contribute Delete
5.53 kB
"""
src/training.py
Training pipeline for FirenetCNN model
"""
import tensorflow as tf
from pathlib import Path
from typing import Optional, Tuple, Dict, Any
from .model import FireNetModel
def train(
train_dir: str,
val_dir: str,
model_save_path: str = 'models/FirenetCNN.keras',
epochs: int = 100,
learning_rate: float = 0.0001,
batch_size: int = 32,
image_size: Tuple[int, int] = (224, 224),
fine_tune: bool = False,
fine_tune_epochs: int = 10,
fine_tune_lr: float = 1e-5,
callback_save_best: bool = True,
callback_early_stopping: bool = True,
early_stopping_patience: int = 10,
) -> Dict[str, Any]:
"""
Train the FirenetCNN model.
Args:
train_dir: Path to training data directory
val_dir: Path to validation data directory
model_save_path: Path to save the trained model
epochs: Number of training epochs
learning_rate: Learning rate for initial training
batch_size: Training batch size
image_size: Input image size (height, width)
fine_tune: Whether to fine-tune the base model after initial training
fine_tune_epochs: Number of fine-tuning epochs
fine_tune_lr: Learning rate for fine-tuning
callback_save_best: Save model with best validation accuracy
callback_early_stopping: Stop training if validation loss stops improving
early_stopping_patience: Patience for early stopping
Returns:
Dictionary with training history and results
"""
# Build model
model = FireNetModel.build_model(input_shape=(*image_size, 3))
# Create data generators
train_gen, val_gen = FireNetModel.create_data_generators(train_dir, val_dir)
# Setup callbacks
callbacks = []
save_path = Path(model_save_path)
save_path.parent.mkdir(parents=True, exist_ok=True)
if callback_save_best:
best_model_path = save_path.with_name(save_path.stem + '_best' + save_path.suffix)
callbacks.append(tf.keras.callbacks.ModelCheckpoint(
str(best_model_path),
monitor='val_accuracy',
save_best_only=True,
mode='max',
verbose=1
))
if callback_early_stopping:
callbacks.append(tf.keras.callbacks.EarlyStopping(
monitor='val_loss',
patience=early_stopping_patience,
restore_best_weights=True,
verbose=1
))
# Initial training (frozen base model)
print(f"Phase 1: Training with frozen base for {epochs} epochs...")
history = model.fit(
train_gen,
epochs=epochs,
validation_data=val_gen,
callbacks=callbacks,
verbose=1
)
# Fine-tuning phase
if fine_tune:
print(f"\nPhase 2: Fine-tuning for {fine_tune_epochs} epochs...")
# Unfreeze the base model
base_model = None
for layer in model.layers:
if hasattr(layer, 'layers') and len(layer.layers) > 50:
base_model = layer
break
if base_model is not None:
base_model.trainable = True
# Recompile with lower learning rate
model.compile(
optimizer=tf.keras.optimizers.Adam(learning_rate=fine_tune_lr),
loss='categorical_crossentropy',
metrics=['accuracy']
)
# Continue training
history_fine = model.fit(
train_gen,
epochs=fine_tune_epochs,
validation_data=val_gen,
callbacks=callbacks,
verbose=1
)
# Merge histories
for key in history_fine.history:
history.history[key].extend(history_fine.history[key])
# Save final model
model.save(str(save_path))
print(f"\nModel saved to: {save_path}")
return {
'history': history.history,
'epochs_completed': len(history.history['accuracy']),
'final_train_acc': history.history['accuracy'][-1],
'final_val_acc': history.history['val_accuracy'][-1],
'model_path': str(save_path)
}
if __name__ == '__main__':
import argparse
parser = argparse.ArgumentParser(description='Train FirenetCNN model')
parser.add_argument('--train-dir', default='data/forestfire-classifier-dataset/train',
help='Training data directory')
parser.add_argument('--val-dir', default='data/forestfire-classifier-dataset/val',
help='Validation data directory')
parser.add_argument('--model-path', default='models/FirenetCNN.keras',
help='Path to save trained model')
parser.add_argument('--epochs', type=int, default=100,
help='Number of training epochs')
parser.add_argument('--fine-tune', action='store_true',
help='Enable fine-tuning phase')
parser.add_argument('--fine-tune-epochs', type=int, default=10,
help='Number of fine-tuning epochs')
args = parser.parse_args()
results = train(
train_dir=args.train_dir,
val_dir=args.val_dir,
model_save_path=args.model_path,
epochs=args.epochs,
fine_tune=args.fine_tune,
fine_tune_epochs=args.fine_tune_epochs
)
print(f"\nTraining complete!")
print(f"Final train accuracy: {results['final_train_acc']:.4f}")
print(f"Final val accuracy: {results['final_val_acc']:.4f}")