Commit ·
25cb1ee
1
Parent(s): 8c47d48
Test
Browse files- dataset_lstm.py +31 -3
- distill_bert_to_lstm.py +8 -2
- example_uses.md +1 -1
- inference_lstm.py +18 -10
dataset_lstm.py
CHANGED
|
@@ -76,6 +76,29 @@ class LSTMTokenizer:
|
|
| 76 |
'input_ids': torch.tensor(ids, dtype=torch.long),
|
| 77 |
'attention_mask': torch.tensor(attention_mask, dtype=torch.long)
|
| 78 |
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 79 |
|
| 80 |
class LSTMDataset(Dataset):
|
| 81 |
"""Dataset for LSTM model"""
|
|
@@ -109,7 +132,7 @@ class LSTMDataset(Dataset):
|
|
| 109 |
|
| 110 |
def prepare_lstm_data(data_path, text_col='text', label_col='label',
|
| 111 |
max_vocab_size=30000, max_seq_length=512,
|
| 112 |
-
val_split=0.1, test_split=0.1, batch_size=32, seed=42, return_datasets=False):
|
| 113 |
"""
|
| 114 |
Load data and prepare for LSTM model
|
| 115 |
"""
|
|
@@ -162,8 +185,10 @@ def prepare_lstm_data(data_path, text_col='text', label_col='label',
|
|
| 162 |
val_dataset = LSTMDataset(val_texts, val_labels, tokenizer)
|
| 163 |
test_dataset = LSTMDataset(test_texts, test_labels, tokenizer)
|
| 164 |
|
| 165 |
-
if return_datasets:
|
| 166 |
return train_dataset, val_dataset, test_dataset, tokenizer.vocab_size
|
|
|
|
|
|
|
| 167 |
|
| 168 |
# Create data loaders
|
| 169 |
if len(train_dataset.texts) == 0:
|
|
@@ -184,4 +209,7 @@ def prepare_lstm_data(data_path, text_col='text', label_col='label',
|
|
| 184 |
else:
|
| 185 |
test_loader = DataLoader(test_dataset, batch_size=batch_size)
|
| 186 |
|
| 187 |
-
|
|
|
|
|
|
|
|
|
|
|
|
| 76 |
'input_ids': torch.tensor(ids, dtype=torch.long),
|
| 77 |
'attention_mask': torch.tensor(attention_mask, dtype=torch.long)
|
| 78 |
}
|
| 79 |
+
|
| 80 |
+
def from_json(self, json_file):
|
| 81 |
+
"""Load tokenizer from JSON file container word2idx dict"""
|
| 82 |
+
import json
|
| 83 |
+
with open(json_file, 'r') as f:
|
| 84 |
+
data = json.load(f)
|
| 85 |
+
self.word2idx = data
|
| 86 |
+
self.idx2word = {v: k for k, v in self.word2idx.items()}
|
| 87 |
+
# Add pad and unk tokens
|
| 88 |
+
self.word2idx['<pad>'] = 0
|
| 89 |
+
self.word2idx['<unk>'] = 1
|
| 90 |
+
self.idx2word[0] = '<pad>'
|
| 91 |
+
self.idx2word[1] = '<unk>'
|
| 92 |
+
# Update vocab size
|
| 93 |
+
self.vocab_size = len(self.word2idx)
|
| 94 |
+
logger.info(f"Loaded tokenizer with {self.vocab_size} tokens")
|
| 95 |
+
|
| 96 |
+
def save(self, json_file):
|
| 97 |
+
"""Save tokenizer word2idx to JSON file"""
|
| 98 |
+
import json
|
| 99 |
+
with open(json_file, 'w') as f:
|
| 100 |
+
json.dump(self.word2idx, f, ensure_ascii=False, indent=4)
|
| 101 |
+
logger.info(f"Tokenizer saved to {json_file}")
|
| 102 |
|
| 103 |
class LSTMDataset(Dataset):
|
| 104 |
"""Dataset for LSTM model"""
|
|
|
|
| 132 |
|
| 133 |
def prepare_lstm_data(data_path, text_col='text', label_col='label',
|
| 134 |
max_vocab_size=30000, max_seq_length=512,
|
| 135 |
+
val_split=0.1, test_split=0.1, batch_size=32, seed=42, return_datasets=False, return_tokenizer=False):
|
| 136 |
"""
|
| 137 |
Load data and prepare for LSTM model
|
| 138 |
"""
|
|
|
|
| 185 |
val_dataset = LSTMDataset(val_texts, val_labels, tokenizer)
|
| 186 |
test_dataset = LSTMDataset(test_texts, test_labels, tokenizer)
|
| 187 |
|
| 188 |
+
if return_datasets and not return_tokenizer:
|
| 189 |
return train_dataset, val_dataset, test_dataset, tokenizer.vocab_size
|
| 190 |
+
elif return_tokenizer and not return_datasets:
|
| 191 |
+
return train_dataset, val_dataset, test_dataset, tokenizer
|
| 192 |
|
| 193 |
# Create data loaders
|
| 194 |
if len(train_dataset.texts) == 0:
|
|
|
|
| 209 |
else:
|
| 210 |
test_loader = DataLoader(test_dataset, batch_size=batch_size)
|
| 211 |
|
| 212 |
+
if not return_tokenizer:
|
| 213 |
+
return train_loader, val_loader, test_loader, tokenizer.vocab_size
|
| 214 |
+
else:
|
| 215 |
+
return train_loader, val_loader, test_loader, tokenizer
|
distill_bert_to_lstm.py
CHANGED
|
@@ -3,6 +3,7 @@ import os
|
|
| 3 |
import logging
|
| 4 |
import torch
|
| 5 |
import random
|
|
|
|
| 6 |
import numpy as np
|
| 7 |
from model import DocBERT
|
| 8 |
from models.lstm_model import DocumentBiLSTM
|
|
@@ -127,15 +128,16 @@ def main():
|
|
| 127 |
|
| 128 |
# Create LSTM data loaders
|
| 129 |
logger.info("Creating LSTM data loaders...")
|
| 130 |
-
lstm_train_loader, lstm_val_loader, lstm_test_loader,
|
| 131 |
args.data_path,
|
| 132 |
text_col=args.text_column,
|
| 133 |
label_col=args.label_column,
|
| 134 |
max_vocab_size=30000,
|
| 135 |
max_seq_length=args.max_seq_length,
|
| 136 |
batch_size=args.batch_size,
|
| 137 |
-
seed=args.seed
|
| 138 |
)
|
|
|
|
| 139 |
|
| 140 |
logger.info(f"LSTM Vocabulary size: {vocab_size}")
|
| 141 |
|
|
@@ -187,6 +189,10 @@ def main():
|
|
| 187 |
save_path = os.path.join(args.output_dir, "distilled_lstm_model.pth")
|
| 188 |
trainer.train(epochs=args.epochs, save_path=save_path)
|
| 189 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 190 |
logger.info("Knowledge distillation completed!")
|
| 191 |
|
| 192 |
if __name__ == "__main__":
|
|
|
|
| 3 |
import logging
|
| 4 |
import torch
|
| 5 |
import random
|
| 6 |
+
import json
|
| 7 |
import numpy as np
|
| 8 |
from model import DocBERT
|
| 9 |
from models.lstm_model import DocumentBiLSTM
|
|
|
|
| 128 |
|
| 129 |
# Create LSTM data loaders
|
| 130 |
logger.info("Creating LSTM data loaders...")
|
| 131 |
+
lstm_train_loader, lstm_val_loader, lstm_test_loader, tokenizer = prepare_lstm_data(
|
| 132 |
args.data_path,
|
| 133 |
text_col=args.text_column,
|
| 134 |
label_col=args.label_column,
|
| 135 |
max_vocab_size=30000,
|
| 136 |
max_seq_length=args.max_seq_length,
|
| 137 |
batch_size=args.batch_size,
|
| 138 |
+
seed=args.seed, return_tokenizer=True
|
| 139 |
)
|
| 140 |
+
vocab_size = tokenizer.vocab_size
|
| 141 |
|
| 142 |
logger.info(f"LSTM Vocabulary size: {vocab_size}")
|
| 143 |
|
|
|
|
| 189 |
save_path = os.path.join(args.output_dir, "distilled_lstm_model.pth")
|
| 190 |
trainer.train(epochs=args.epochs, save_path=save_path)
|
| 191 |
|
| 192 |
+
# Save the tokenizer
|
| 193 |
+
tokenizer_path = os.path.join(args.output_dir, "tokenizer.json")
|
| 194 |
+
tokenizer.save(tokenizer_path)
|
| 195 |
+
|
| 196 |
logger.info("Knowledge distillation completed!")
|
| 197 |
|
| 198 |
if __name__ == "__main__":
|
example_uses.md
CHANGED
|
@@ -17,5 +17,5 @@ python ./distill_bert_to_lstm.py --bert_model bert-base-uncased --bert_model_pat
|
|
| 17 |
|
| 18 |
- Inference with distilled LSTM model (test_data.csv is test dataset with 4 classes like ag_news)
|
| 19 |
```
|
| 20 |
-
python ./inference_lstm.py --model_path "./docbert_lstm/distilled_lstm_model.pth" --num_classes 4 --class_names "World" "Sports" "Business" "Science" --text_column "Description" --label_column "Class Index" --data_path "./test_data.csv" --inference_batch_limit 10
|
| 21 |
```
|
|
|
|
| 17 |
|
| 18 |
- Inference with distilled LSTM model (test_data.csv is test dataset with 4 classes like ag_news)
|
| 19 |
```
|
| 20 |
+
python ./inference_lstm.py --model_path "./docbert_lstm/distilled_lstm_model.pth" --num_classes 4 --class_names "World" "Sports" "Business" "Science" --text_column "Description" --label_column "Class Index" --data_path "./test_data.csv" --inference_batch_limit 10 --tokenizer_path "./docbert_lstm/tokenizer.json"
|
| 21 |
```
|
inference_lstm.py
CHANGED
|
@@ -1,7 +1,7 @@
|
|
| 1 |
from dataset_lstm import prepare_lstm_data, LSTMTokenizer, LSTMDataset
|
| 2 |
from models.lstm_model import DocumentBiLSTM
|
| 3 |
from sklearn import metrics
|
| 4 |
-
import torch
|
| 5 |
import numpy as np
|
| 6 |
import argparse
|
| 7 |
|
|
@@ -9,6 +9,7 @@ if __name__ == "__main__":
|
|
| 9 |
parser = argparse.ArgumentParser(description="Document Classification with LSTM")
|
| 10 |
parser.add_argument("--data_path", type=str, required=True, help="Path to the dataset")
|
| 11 |
parser.add_argument("--model_path", type=str, required=True, help="Path to the trained model")
|
|
|
|
| 12 |
parser.add_argument("--max_seq_length", type=int, default=512, help="Maximum sequence length for LSTM")
|
| 13 |
parser.add_argument("--batch_size", type=int, default=32, help="Batch size for training and evaluation")
|
| 14 |
parser.add_argument("--num_classes", type=int, required=True, help="Number of classes for classification")
|
|
@@ -30,13 +31,26 @@ if __name__ == "__main__":
|
|
| 30 |
# Set device
|
| 31 |
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 32 |
|
| 33 |
-
|
|
|
|
| 34 |
tokenizer = LSTMTokenizer(max_seq_length=args.max_seq_length)
|
| 35 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 36 |
text_col=args.text_column,
|
| 37 |
label_col=args.label_column,
|
| 38 |
batch_size=args.batch_size,
|
| 39 |
-
max_seq_length=args.max_seq_length, val_split=0.0, test_split=
|
| 40 |
|
| 41 |
test_loader, _, _, vocab_size = prepare_lstm_data(args.data_path,
|
| 42 |
text_col=args.text_column,
|
|
@@ -50,8 +64,6 @@ if __name__ == "__main__":
|
|
| 50 |
hidden_dim=args.hidden_dim,
|
| 51 |
n_layers=args.num_layers,
|
| 52 |
output_dim=args.num_classes)
|
| 53 |
-
|
| 54 |
-
model_state = torch.load(args.model_path)
|
| 55 |
|
| 56 |
if 'vocab_size' in model_state:
|
| 57 |
vocab_size = model_state['vocab_size']
|
|
@@ -77,10 +89,6 @@ if __name__ == "__main__":
|
|
| 77 |
|
| 78 |
outputs = model(input_ids)
|
| 79 |
|
| 80 |
-
if batch_count == 0 or batch_count == 1:
|
| 81 |
-
print(f"Labels: {labels}")
|
| 82 |
-
print(f"Outputs: {outputs}")
|
| 83 |
-
|
| 84 |
predictions = torch.argmax(outputs, dim=1)
|
| 85 |
all_predictions = np.append(all_predictions, predictions.cpu().numpy())
|
| 86 |
|
|
|
|
| 1 |
from dataset_lstm import prepare_lstm_data, LSTMTokenizer, LSTMDataset
|
| 2 |
from models.lstm_model import DocumentBiLSTM
|
| 3 |
from sklearn import metrics
|
| 4 |
+
import torch, logging
|
| 5 |
import numpy as np
|
| 6 |
import argparse
|
| 7 |
|
|
|
|
| 9 |
parser = argparse.ArgumentParser(description="Document Classification with LSTM")
|
| 10 |
parser.add_argument("--data_path", type=str, required=True, help="Path to the dataset")
|
| 11 |
parser.add_argument("--model_path", type=str, required=True, help="Path to the trained model")
|
| 12 |
+
parser.add_argument("--tokenizer_path", type=str, required=True, help="Path to the tokenizer")
|
| 13 |
parser.add_argument("--max_seq_length", type=int, default=512, help="Maximum sequence length for LSTM")
|
| 14 |
parser.add_argument("--batch_size", type=int, default=32, help="Batch size for training and evaluation")
|
| 15 |
parser.add_argument("--num_classes", type=int, required=True, help="Number of classes for classification")
|
|
|
|
| 31 |
# Set device
|
| 32 |
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 33 |
|
| 34 |
+
model_state = torch.load(args.model_path)
|
| 35 |
+
|
| 36 |
tokenizer = LSTMTokenizer(max_seq_length=args.max_seq_length)
|
| 37 |
+
#tokenizer.from_json(args.tokenizer_path)
|
| 38 |
+
|
| 39 |
+
_, _, _, tokenizer = prepare_lstm_data("./test_data.csv",
|
| 40 |
+
text_col=args.text_column,
|
| 41 |
+
label_col=args.label_column,
|
| 42 |
+
batch_size=args.batch_size,
|
| 43 |
+
max_seq_length=args.max_seq_length, val_split=0.0, test_split=0.0, return_datasets=False, return_tokenizer=True)
|
| 44 |
+
|
| 45 |
+
tokenizer.save(args.tokenizer_path)
|
| 46 |
+
|
| 47 |
+
# Prepare data
|
| 48 |
+
|
| 49 |
+
_, _, test_dataset, vocab_size = prepare_lstm_data(args.data_path,
|
| 50 |
text_col=args.text_column,
|
| 51 |
label_col=args.label_column,
|
| 52 |
batch_size=args.batch_size,
|
| 53 |
+
max_seq_length=args.max_seq_length, val_split=0.0, test_split=1.0, return_datasets=True)
|
| 54 |
|
| 55 |
test_loader, _, _, vocab_size = prepare_lstm_data(args.data_path,
|
| 56 |
text_col=args.text_column,
|
|
|
|
| 64 |
hidden_dim=args.hidden_dim,
|
| 65 |
n_layers=args.num_layers,
|
| 66 |
output_dim=args.num_classes)
|
|
|
|
|
|
|
| 67 |
|
| 68 |
if 'vocab_size' in model_state:
|
| 69 |
vocab_size = model_state['vocab_size']
|
|
|
|
| 89 |
|
| 90 |
outputs = model(input_ids)
|
| 91 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 92 |
predictions = torch.argmax(outputs, dim=1)
|
| 93 |
all_predictions = np.append(all_predictions, predictions.cpu().numpy())
|
| 94 |
|