CoolFace
Apppublic

YanushanthR/URLs_Classifier_with_Self_Learning

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
retrain_model.py88 linesDownload Raw Back to root
1import pandas as pd2from sklearn.ensemble import RandomForestClassifier3from sklearn.preprocessing import OneHotEncoder4import joblib5 6# Paths for dataset and model7dataset_path = "dataset.csv"8model_path = "RF_trained_model.pkl"9feature_names_path = "feature_names.pkl"10 11# Load dataset12try:13    data = pd.read_csv(dataset_path)14    if data.empty:15        raise ValueError("Dataset is empty. Cannot retrain the model.")16except Exception as e:17    print(f"Error loading dataset: {e}")18    exit(1)19 20# Feature extraction function21def extract_features(url):22    from urllib.parse import urlparse23    from math import log224    import re25 26    parsed_url = urlparse(url)27    path = parsed_url.path28    query = parsed_url.query29 30    # Entropy calculation31    probabilities = [float(url.count(c)) / len(url) for c in set(url)]32    url_entropy = -sum(p * log2(p) for p in probabilities) if len(url) > 0 else 033 34    # Helper functions35    contains_ip = lambda url: int(bool(re.search(r"(?:\d{1,3}\.){3}\d{1,3}", url)))36    has_encoded_chars = lambda url: int("%" in url)37    has_suspicious_substrings = lambda url: int(any(sub in url for sub in ["@", "-", "_", "~"]))38    has_malicious_file_extension = lambda url: int(any(url.lower().endswith(ext) for ext in [".exe", ".zip", ".js", ".rar", ".bat"]))39 40    return {41        "url_length": len(url),42        "num_digits": sum(c.isdigit() for c in url),43        "num_special_chars": sum(not c.isalnum() for c in url),44        "has_https": int("https" in url.lower()),45        "num_subdomains": url.count('.'),46        "has_suspicious_keywords": int(any(kw in url.lower() for kw in ["login", "secure", "account", "update"])),47        "url_entropy": url_entropy,48        "domain_length": len(parsed_url.netloc),49        "path_length": len(path),50        "presence_of_ip": contains_ip(url),51        "tld": parsed_url.netloc.split('.')[-1] if '.' in parsed_url.netloc else "",52        "num_query_params": len(query.split('&')) if query else 0,53        "has_encoded_chars": has_encoded_chars(url),54        "path_depth": len(path.split('/')) - 1 if path else 0,55        "has_suspicious_substrings": has_suspicious_substrings(url),56        "has_malicious_file_extension": has_malicious_file_extension(url)57    }58 59# Extract features and labels60def extract_features_and_labels(data):61    features = [extract_features(url) for url in data["url"]]62    X = pd.DataFrame(features)63    y = data["label"]64    return X, y65 66try:67    # Extract features and labels68    X, y = extract_features_and_labels(data)69 70    # One-hot encode the 'tld' column71    if "tld" in X.columns:72        encoder = OneHotEncoder(sparse=False, handle_unknown="ignore")73        tld_encoded = encoder.fit_transform(X[["tld"]])74        tld_encoded_df = pd.DataFrame(tld_encoded, columns=encoder.get_feature_names_out(["tld"]))75        X = pd.concat([X.drop(columns=["tld"]), tld_encoded_df], axis=1)76 77    # Train the Random Forest model78    model = RandomForestClassifier()79    model.fit(X, y)80 81    # Save the model and feature names82    joblib.dump(model, model_path)83    joblib.dump(X.columns.tolist(), feature_names_path)84    print("Model retrained and saved successfully.")85except Exception as e:86    print(f"Error during retraining: {e}")87    exit(2)88