jesse-tong commited on
Commit
25cb1ee
·
1 Parent(s): 8c47d48
Files changed (4) hide show
  1. dataset_lstm.py +31 -3
  2. distill_bert_to_lstm.py +8 -2
  3. example_uses.md +1 -1
  4. 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
- return train_loader, val_loader, test_loader, tokenizer.vocab_size
 
 
 
 
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, vocab_size = prepare_lstm_data(
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
- # Prepare data
 
34
  tokenizer = LSTMTokenizer(max_seq_length=args.max_seq_length)
35
- test_dataset, _, _, vocab_size = prepare_lstm_data(args.data_path,
 
 
 
 
 
 
 
 
 
 
 
 
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=0.0, return_datasets=True)
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