lyzq/malicious-url-app
0
1from __future__ import annotations2 3from collections import Counter4from copy import deepcopy5 6import numpy as np7import torch8from sklearn.base import BaseEstimator, ClassifierMixin9from sklearn.model_selection import train_test_split10from torch import nn11from torch.utils.data import DataLoader, TensorDataset12 13from src.config import RANDOM_STATE14 15 16class CharCNN(nn.Module):17 def __init__(self, vocab_size: int, embedding_dim: int, channels: int, dropout: float) -> None:18 super().__init__()19 self.embedding = nn.Embedding(vocab_size, embedding_dim, padding_idx=0)20 self.conv3 = nn.Conv1d(embedding_dim, channels, kernel_size=3, padding=1)21 self.conv5 = nn.Conv1d(embedding_dim, channels, kernel_size=5, padding=2)22 self.activation = nn.ReLU()23 self.dropout = nn.Dropout(dropout)24 self.classifier = nn.Sequential(25 nn.Linear(channels * 2, channels),26 nn.ReLU(),27 nn.Dropout(dropout),28 nn.Linear(channels, 1),29 )30 31 def forward(self, inputs: torch.Tensor) -> torch.Tensor:32 embedded = self.embedding(inputs).transpose(1, 2)33 conv3 = self.activation(self.conv3(embedded))34 conv5 = self.activation(self.conv5(embedded))35 pooled3 = torch.amax(conv3, dim=2)36 pooled5 = torch.amax(conv5, dim=2)37 features = self.dropout(torch.cat([pooled3, pooled5], dim=1))38 return self.classifier(features).squeeze(1)39 40 41class CharCNNURLClassifier(BaseEstimator, ClassifierMixin):42 def __init__(43 self,44 max_length: int = 200,45 max_vocab_size: int = 128,46 embedding_dim: int = 32,47 channels: int = 64,48 dropout: float = 0.25,49 batch_size: int = 512,50 epochs: int = 4,51 learning_rate: float = 1e-3,52 validation_size: float = 0.1,53 ) -> None:54 self.max_length = max_length55 self.max_vocab_size = max_vocab_size56 self.embedding_dim = embedding_dim57 self.channels = channels58 self.dropout = dropout59 self.batch_size = batch_size60 self.epochs = epochs61 self.learning_rate = learning_rate62 self.validation_size = validation_size63 64 def _build_vocab(self, urls: list[str]) -> dict[str, int]:65 counter = Counter()66 for url in urls:67 counter.update(url[: self.max_length])68 most_common = [char for char, _ in counter.most_common(self.max_vocab_size - 2)]69 vocab = {"<pad>": 0, "<unk>": 1}70 for idx, char in enumerate(most_common, start=2):71 vocab[char] = idx72 return vocab73 74 def _encode_urls(self, urls: list[str]) -> np.ndarray:75 encoded = np.zeros((len(urls), self.max_length), dtype=np.int64)76 for row_idx, url in enumerate(urls):77 for col_idx, char in enumerate(url[: self.max_length]):78 encoded[row_idx, col_idx] = self.vocab_.get(char, 1)79 return encoded80 81 def _build_model(self) -> CharCNN:82 return CharCNN(83 vocab_size=len(self.vocab_),84 embedding_dim=self.embedding_dim,85 channels=self.channels,86 dropout=self.dropout,87 )88 89 def fit(self, X, y):90 urls = [str(item) for item in X]91 labels = np.asarray(y, dtype=np.float32)92 self.vocab_ = self._build_vocab(urls)93 self.classes_ = np.array([0, 1])94 95 encoded = self._encode_urls(urls)96 X_train, X_val, y_train, y_val = train_test_split(97 encoded,98 labels,99 test_size=self.validation_size,100 stratify=labels,101 random_state=RANDOM_STATE,102 )103 104 self.device_ = "cuda" if torch.cuda.is_available() else "cpu"105 self.model_ = self._build_model().to(self.device_)106 107 train_dataset = TensorDataset(108 torch.tensor(X_train, dtype=torch.long),109 torch.tensor(y_train, dtype=torch.float32),110 )111 val_dataset = TensorDataset(112 torch.tensor(X_val, dtype=torch.long),113 torch.tensor(y_val, dtype=torch.float32),114 )115 116 train_loader = DataLoader(train_dataset, batch_size=self.batch_size, shuffle=True)117 val_loader = DataLoader(val_dataset, batch_size=self.batch_size, shuffle=False)118 119 positive_count = float(y_train.sum())120 negative_count = float(len(y_train) - positive_count)121 pos_weight = torch.tensor([negative_count / max(positive_count, 1.0)], device=self.device_)122 123 criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight)124 optimizer = torch.optim.AdamW(self.model_.parameters(), lr=self.learning_rate)125 126 best_state = deepcopy(self.model_.state_dict())127 best_val_loss = float("inf")128 self.history_ = []129 130 for epoch in range(self.epochs):131 self.model_.train()132 train_loss = 0.0133 for batch_inputs, batch_targets in train_loader:134 batch_inputs = batch_inputs.to(self.device_)135 batch_targets = batch_targets.to(self.device_)136 optimizer.zero_grad()137 logits = self.model_(batch_inputs)138 loss = criterion(logits, batch_targets)139 loss.backward()140 optimizer.step()141 train_loss += loss.item() * len(batch_targets)142 143 self.model_.eval()144 val_loss = 0.0145 with torch.no_grad():146 for batch_inputs, batch_targets in val_loader:147 batch_inputs = batch_inputs.to(self.device_)148 batch_targets = batch_targets.to(self.device_)149 logits = self.model_(batch_inputs)150 loss = criterion(logits, batch_targets)151 val_loss += loss.item() * len(batch_targets)152 153 avg_train_loss = train_loss / len(train_dataset)154 avg_val_loss = val_loss / len(val_dataset)155 self.history_.append(156 {"epoch": epoch + 1, "train_loss": float(avg_train_loss), "val_loss": float(avg_val_loss)}157 )158 if avg_val_loss < best_val_loss:159 best_val_loss = avg_val_loss160 best_state = deepcopy(self.model_.state_dict())161 162 self.model_.load_state_dict(best_state)163 self.model_.to("cpu")164 self.device_ = "cpu"165 self.model_.eval()166 self.best_val_loss_ = float(best_val_loss)167 return self168 169 def predict_proba(self, X):170 urls = [str(item) for item in X]171 encoded = self._encode_urls(urls)172 dataset = TensorDataset(torch.tensor(encoded, dtype=torch.long))173 loader = DataLoader(dataset, batch_size=self.batch_size, shuffle=False)174 175 self.model_.eval()176 scores: list[np.ndarray] = []177 with torch.no_grad():178 for (batch_inputs,) in loader:179 logits = self.model_(batch_inputs)180 probs = torch.sigmoid(logits).cpu().numpy()181 scores.append(probs)182 positive_probs = np.concatenate(scores)183 return np.column_stack([1.0 - positive_probs, positive_probs])184 185 def predict(self, X):186 return (self.predict_proba(X)[:, 1] >= 0.5).astype(int)187 188 def __getstate__(self):189 state = self.__dict__.copy()190 if "model_" in state:191 state["model_state_dict_"] = {key: value.cpu() for key, value in self.model_.state_dict().items()}192 del state["model_"]193 return state194 195 def __setstate__(self, state):196 self.__dict__.update(state)197 if "model_state_dict_" in state and "vocab_" in state:198 self.model_ = self._build_model()199 self.model_.load_state_dict(state["model_state_dict_"])200 self.model_.to("cpu")201 self.model_.eval()202 203 