File size: 7,177 Bytes
77bc910 acf97d8 8bac2cf 0ab884f 6e638f9 0ab884f acf97d8 63812ec acf97d8 dc5d564 acf97d8 f292cd1 acf97d8 01b0ca9 acf97d8 83f2f41 acf97d8 77bc910 f292cd1 a4d7cd8 77bc910 f292cd1 77bc910 a4d7cd8 77bc910 a4d7cd8 77bc910 63812ec 306c7ca a4d7cd8 acf97d8 831ff7d 01b0ca9 e416d63 f292cd1 ca56e1c 4f44808 8a6c918 4f44808 626f169 be59bb2 626f169 be59bb2 626f169 acf97d8 efb13cd acf97d8 2117892 acf97d8 2117892 6e638f9 f292cd1 6e638f9 44c78a2 0ea4a8b acf97d8 f292cd1 acf97d8 8bac2cf acf97d8 | 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 | from dataset import load_data, create_data_loaders
from models.lstm_model import DocumentBiLSTM
from sklearn import metrics
import torch, random
import torch.nn.functional as F
from torch.utils.data import DataLoader
import numpy as np
import argparse
# Add these imports for mapping optimization
from itertools import permutations
import copy
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Document Classification with LSTM")
parser.add_argument("--data_path", type=str, required=True, help="Path to the dataset")
parser.add_argument("--bert_model", type=str, default="bert-base-uncased", help="BERT model name or path used for distillation (as we'll use its tokenizer)")
parser.add_argument("--model_path", type=str, required=True, help="Path to the trained model")
parser.add_argument("--max_seq_length", type=int, default=512, help="Maximum sequence length for LSTM")
parser.add_argument("--batch_size", type=int, default=32, help="Batch size for training and evaluation")
parser.add_argument("--num_classes", type=int, required=True, help="Number of classes for classification")
parser.add_argument("--text_column", type=str, default="text", help="Column name for text data")
parser.add_argument("--label_column", type=str, nargs='+', help="Column name for labels")
parser.add_argument("--class_names", type=str, nargs='+', required=True, help="List of class names for classification")
parser.add_argument("--inference_batch_limit", type=int, default=-1, help="Limit for inference batch counts")
parser.add_argument("--print_predictions", type=bool, default=False, help="Print predictions to console")
# 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")
args = parser.parse_args()
class_names = args.class_names
# Set device
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model_state = torch.load(args.model_path, map_location=device)
# 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, val_data, test_data = load_data(
args.data_path,
text_col=args.text_column,
label_col=label_column,
validation_split=0.0,
test_split=1.0,
seed=42
)
# Create BERT data loaders
print("Creating data loaders (note the datasets and dataloaders use BERT's tokenizer)...")
train_dataset, val_dataset, 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
)
bert_vocab_size = train_dataset.tokenizer.vocab_size
test_loader = DataLoader(test_dataset, batch_size=args.batch_size, shuffle=False)
# Load model
model = DocumentBiLSTM(vocab_size=bert_vocab_size,
embedding_dim=args.embedding_dim,
hidden_dim=args.hidden_dim,
n_layers=args.num_layers,
output_dim=args.num_classes * num_categories)
# I don't know why the model is trained with 30000 embedding size (maybe I forgot to update the distillation code before training)
# so this is a temporary fix
if model_state['model_state_dict']['embedding.weight'].shape[0] == 30000:
model = DocumentBiLSTM(vocab_size=30000,
embedding_dim=args.embedding_dim,
hidden_dim=args.hidden_dim,
n_layers=args.num_layers,
output_dim=args.num_classes)
if 'model_state_dict' in model_state:
model.load_state_dict(model_state['model_state_dict'], strict=False)
else:
model.load_state_dict(model_state, strict=False)
model = model.to(device)
all_labels = np.array([], dtype=int)
all_predictions = np.array([], dtype=int)
# Inference
batch_count = 0
with torch.no_grad():
for batch in test_loader:
input_ids = batch['input_ids'].to(device)
labels = batch['label'].to(device)
attention_mask = batch['attention_mask'].to(device)
all_labels = np.append(all_labels, labels.cpu().numpy())
outputs = model(input_ids, attention_mask=attention_mask)
probs = F.softmax(outputs, dim=1)
batch_size, total_classes = outputs.shape
if total_classes % num_categories != 0:
raise ValueError(f"Error: Number of total classes in the batch must of divisible by {num_categories}")
classes_per_group = total_classes // num_categories
# Group every classes_per_group values along dim=1
reshaped = outputs.view(outputs.size(0), -1, classes_per_group) # shape: (batch, self., classes_per_group)
# Argmax over each group of classes_per_group
preds = reshaped.argmax(dim=-1)
predictions = torch.argmax(probs, dim=1)
all_predictions = np.append(all_predictions, predictions.cpu().numpy())
if args.print_predictions:
for i in range(len(predictions)):
print(f"Text: {test_dataset.get_text_(batch_count * args.batch_size + i)}, Prediction: {predictions[i]}, True Label: {labels[i]}")
if args.inference_batch_limit > 0 and batch_count >= args.inference_batch_limit:
break
batch_count += 1
# Print classification report
# Calculate accuracy, F1 score, recall, and precision
accuracy = metrics.accuracy_score(all_labels, all_predictions)
f1 = metrics.f1_score(all_labels, all_predictions, average='weighted')
precision = metrics.precision_score(all_labels, all_predictions, average='weighted')
recall = metrics.recall_score(all_labels, all_predictions, average='weighted')
print(f"Accuracy: {accuracy}")
print(f"F1 Score: {f1}")
print(f"Precision: {precision}")
print(f"Recall: {recall}")
with open("predictions_lstm.txt", "w") as f:
for i in range(len(all_labels)):
idx = int(i)
f.write(f"Text: {test_dataset.get_text_(idx)}\n")
f.write(f"True Label: {all_labels[idx]}, Predicted Label: {all_predictions[idx]}\n")
f.write("\n")
with open("metrics_lstm.txt", "w") as f:
f.write(f"Accuracy: {accuracy}\n")
f.write(f"F1 Score: {f1}\n")
f.write(f"Precision: {precision}\n")
f.write(f"Recall: {recall}\n")
|