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()