A simplified BERT finetuning package with embedding learning and classification stages, optimized for single GPU training.
- Masked Language Modeling (MLM)
- Whole Word Masking
- Resource-efficient training
- Mixed precision (FP16)
- Uses MLM-finetuned embeddings
- Hyperparameter optimization
- Resource monitoring
- Gradient checkpointing
- Optuna trials for hyperparameters
- Resource-aware worker scaling
- Memory management
- Early stopping
- Weights & Biases logging
- Resource tracking
- Process monitoring
- Performance metrics
The training process consists of three main stages:
Train the BERT model to learn better embeddings through MLM:
from simpler_fine_bert.embedding.embedding_training import train_embeddings
from simpler_fine_bert.common.config_utils import load_config
from pathlib import Path
config = load_config("config_embedding.yaml")
output_dir = Path("outputs")
loss, metrics = train_embeddings(config, output_dir)Run Optuna trials to find optimal classification hyperparameters:
from simpler_fine_bert.classification.classification_training import run_classification_optimization
from simpler_fine_bert.common.config_utils import load_config
best_params = run_classification_optimization(
embedding_model_path="outputs/embedding_stage/best_model", # Path to your best embedding model
config_path="config_finetune.yaml", # Classification config
study_name="classification_study" # Optional study name
)Train the final classification model using the best parameters:
from simpler_fine_bert.classification.classification_training import train_final_model
from simpler_fine_bert.common.config_utils import load_config
from pathlib import Path
train_final_model(
embedding_model_path="outputs/embedding_stage/best_model", # Path to your best embedding model
best_params=best_params, # Parameters from optimization
config_path="config_finetune.yaml", # Classification config
output_dir=Path("outputs") # Optional output directory
)The package includes several memory optimizations:
- Gradient checkpointing
- Mixed precision (FP16)
- Gradient accumulation
- Resource monitoring
- Memory-aware worker scaling
Each training stage produces:
-
Embedding stage:
- Best MLM model
- Training metrics
- Resource logs
-
Classification optimization:
- Best hyperparameters
- Trial metrics
- Study statistics
-
Final classification:
- Best model checkpoint
- Evaluation metrics
- Performance analysis
This project is licensed under the MIT License - see the LICENSE file for details.