File size: 7,888 Bytes
ae47555 25cb1ee ae47555 4f44808 ae47555 a63576f ae47555 f292cd1 ae47555 8e3d6fe ae47555 f292cd1 a63576f 0227fba ae47555 f292cd1 84fdef2 ae47555 a63576f ae47555 a4d7cd8 ae47555 a4d7cd8 ae47555 a4d7cd8 84fdef2 d252d6b 82406fe a4d7cd8 ae47555 f292cd1 ae47555 82406fe ae47555 82406fe ae47555 f292cd1 ae47555 6cf4c1f ae47555 25cb1ee ae47555 | 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 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 | import argparse
import os
import logging
import torch
import random
import json
import numpy as np
from model import DocBERT
from models.lstm_model import DocumentBiLSTM
from dataset import load_data, create_data_loaders
from knowledge_distillation import DistillationTrainer
from transformers import BertTokenizer
# Setup logging
logging.basicConfig(
format="%(asctime)s - %(levelname)s - %(message)s",
level=logging.INFO,
datefmt="%Y-%m-%d %H:%M:%S",
)
logger = logging.getLogger(__name__)
def set_seed(seed):
"""Set all seeds for reproducibility"""
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
torch.backends.cudnn.deterministic = True
torch.backends.cudnn.benchmark = False
def main():
parser = argparse.ArgumentParser(description="Distill knowledge from BERT to LSTM for document classification")
# Data arguments
parser.add_argument("--train_data_path", type=str, required=True, help="Path to the dataset file (CSV or TSV)")
parser.add_argument("--val_data_path", type=str, required=True, help="Path to the validation dataset file (CSV or TSV)")
parser.add_argument("--test_data_path", type=str, required=True, help="Path to the test dataset file (CSV or TSV)")
parser.add_argument("--text_column", type=str, default="text", help="Name of the text column")
parser.add_argument("--label_column", type=str, nargs="+", help="Name of the label column")
parser.add_argument("--val_split", type=float, default=0.1, help="Validation set split ratio")
parser.add_argument("--test_split", type=float, default=0.1, help="Test set split ratio")
# BERT model arguments
parser.add_argument("--bert_model", type=str, default="bert-base-uncased", help="BERT model to use")
parser.add_argument("--bert_model_path", type=str, required=True, help="Path to saved BERT model weights")
parser.add_argument("--max_seq_length", type=int, default=250, help="Maximum sequence length (e.g., 250 for PhoBERT as PhoBERT allows max_position_embeddings=258)")
# LSTM model arguments
parser.add_argument("--embedding_dim", type=int, default=300, help="Dimension of word embeddings in LSTM")
parser.add_argument("--hidden_dim", type=int, default=256, help="Hidden dimension of LSTM")
parser.add_argument("--num_layers", type=int, default=2, help="Number of LSTM layers")
parser.add_argument("--dropout", type=float, default=0.5, help="Dropout probability")
# Distillation arguments
parser.add_argument("--temperature", type=float, default=2.0, help="Temperature for softening probability distributions")
parser.add_argument("--alpha", type=float, default=0.5, help="Weight for distillation loss vs. regular loss")
parser.add_argument("--num_classes", type=int, required=True, help="Number of classes to predict")
# Training arguments
parser.add_argument("--batch_size", type=int, default=16, help="Training batch size")
parser.add_argument("--learning_rate", type=float, default=0.001, help="Learning rate for LSTM")
parser.add_argument("--epochs", type=int, default=20, help="Number of training epochs")
# Other arguments
parser.add_argument("--seed", type=int, default=42, help="Random seed")
parser.add_argument("--output_dir", type=str, default="./output", help="Directory to save models")
args = parser.parse_args()
# Set seed for reproducibility
set_seed(args.seed)
# Create output directory if it doesn't exist
if not os.path.exists(args.output_dir):
os.makedirs(args.output_dir)
# Load and prepare data for both BERT and LSTM
logger.info("Loading and preparing data...")
# Load data first
label_column = args.label_column[0] if isinstance(args.label_column, list) and len(args.label_column) == 1 else args.label_column
num_categories = len(args.label_column) if isinstance(args.label_column, list) else 1
train_data, _, _ = load_data(
args.train_data_path,
text_col=args.text_column,
label_col=label_column,
validation_split=0.0,
test_split=0.0,
seed=args.seed
)
_, val_data, _ = load_data(
args.val_data_path,
text_col=args.text_column,
label_col=label_column,
validation_split=1.0,
test_split=0.0,
seed=args.seed
)
_, _, test_data = load_data(
args.test_data_path,
text_col=args.text_column,
label_col=label_column,
validation_split=0.0,
test_split=1.0,
seed=args.seed
)
# Create BERT data loaders
logger.info("Creating BERT data loaders...")
bert_train_dataset, bert_val_dataset, bert_test_dataset = create_data_loaders(
train_data,
val_data,
test_data,
tokenizer_name=args.bert_model,
max_length=args.max_seq_length,
batch_size=args.batch_size,
num_classes=args.num_classes,
return_datasets=True
)
print("Train samples: ", len(bert_train_dataset))
print("Validation samples: ", len(bert_val_dataset))
print("Test samples: ", len(bert_test_dataset))
# Create dataloaders
bert_train_loader = torch.utils.data.DataLoader(bert_train_dataset, batch_size=args.batch_size, shuffle=True)
bert_val_loader = torch.utils.data.DataLoader(bert_val_dataset, batch_size=args.batch_size)
bert_test_loader = torch.utils.data.DataLoader(bert_test_dataset, batch_size=args.batch_size)
# Load pre-trained BERT model (teacher)
logger.info("Loading pre-trained BERT model (teacher)...")
bert_model = DocBERT(
num_classes=args.num_classes,
bert_model_name=args.bert_model,
dropout_prob=0.1,
num_categories=num_categories
)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# Load saved BERT weights
bert_model.load_state_dict(torch.load(args.bert_model_path, map_location=device))
logger.info(f"Loaded teacher model from {args.bert_model_path}")
vocab_size = bert_model.bert.tokenizer.vocab_size
logger.info(f"LSTM Vocabulary size: {vocab_size}")
print("LSTM Vocabulary size: ", vocab_size)
# Initialize LSTM model (student)
logger.info("Initializing LSTM model (student)...")
lstm_model = DocumentBiLSTM(
vocab_size=vocab_size,
embedding_dim=args.embedding_dim,
hidden_dim=args.hidden_dim,
output_dim=args.num_classes * num_categories,
n_layers=args.num_layers,
dropout=args.dropout
)
# Print model sizes for comparison
bert_params = sum(p.numel() for p in bert_model.parameters())
lstm_params = sum(p.numel() for p in lstm_model.parameters())
logger.info(f"BERT model size: {bert_params:,} parameters")
logger.info(f"LSTM model size: {lstm_params:,} parameters")
logger.info(f"Size reduction: {bert_params / lstm_params:.1f}x")
# Initialize distillation trainer
trainer = DistillationTrainer(
teacher_model=bert_model,
student_model=lstm_model,
train_loader=bert_train_loader, # Using BERT loader to match tokenization
val_loader=bert_val_loader,
test_loader=bert_test_loader,
temperature=args.temperature,
alpha=args.alpha,
lr=args.learning_rate,
num_categories=num_categories,
num_classes=args.num_classes,
weight_decay=1e-5
)
# Train with knowledge distillation
logger.info("Starting knowledge distillation...")
save_path = os.path.join(args.output_dir, "distilled_lstm_model.pth")
trainer.train(epochs=args.epochs, save_path=save_path)
logger.info("Knowledge distillation completed!")
if __name__ == "__main__":
main() |