Files
AI-Red-Teaming-CSCD94/ai-data-attacks/feature attack/clean-label-student-template.ipynb
T
2026-07-26 22:53:03 -04:00

1.2 MiB

In [1]:
import numpy as np
import matplotlib.pyplot as plt
import json
import requests
from sklearn.linear_model import LogisticRegression
from sklearn.preprocessing import (
    StandardScaler,
)  # May be needed if student wants to scale within function, though data is pre-scaled
from sklearn.multiclass import OneVsRestClassifier
from sklearn.neighbors import NearestNeighbors
from sklearn.metrics import accuracy_score, classification_report
import os

SEED = 1337
np.random.seed(SEED)

# Dataset filename provided by the teacher
dataset_filename = "clean_label_eval_dataset.npz"

# Attack Hyperparameters
N_NEIGHBORS = 12  # Number of neighbors from perturbing class to modify
EPSILON_CROSS = 0.25  # Perturbation magnitude (how far to push neighbors)

htb_green = "#9fef00"
node_black = "#141d2b"
hacker_grey = "#a4b1cd"
white = "#ffffff"
azure = "#0086ff"
nugget_yellow = "#ffaf00"
malware_red = "#ff3e3e"
vivid_purple = "#9f00ff"
aquamarine = "#2ee7b6"

# Configure plot styles
plt.style.use("seaborn-v0_8-darkgrid")
plt.rcParams.update(
    {
        "figure.facecolor": node_black,
        "axes.facecolor": node_black,
        "axes.edgecolor": hacker_grey,
        "axes.labelcolor": white,
        "text.color": white,
        "xtick.color": hacker_grey,
        "ytick.color": hacker_grey,
        "grid.color": hacker_grey,
        "grid.alpha": 0.1,
        "legend.facecolor": node_black,
        "legend.edgecolor": hacker_grey,
        "legend.frameon": True,
        "legend.framealpha": 0.8,
        "legend.labelcolor": white,
        "figure.figsize": (12, 7),
    }
)
In [2]:
try:
    data = np.load(dataset_filename)
    X_train = data["Xtr"]
    y_train = data["ytr"]
    X_test = data["Xte"]
    y_test = data["yte"]
    # Load target index - make sure it's an integer scalar
    target_index = int(data["target_idx"].item())
    data.close()
    print("Data loaded successfully from .npz file.")
    print(f"X_train shape: {X_train.shape}")
    print(f"y_train shape: {y_train.shape}")
    print(f"X_test shape: {X_test.shape}")
    print(f"y_test shape: {y_test.shape}")
    print(f"Target Point Index: {target_index}")

    # Determine target class and infer perturbing/misclassify-as class based on common scenarios
    if target_index < 0 or target_index >= len(y_train):
        raise ValueError(
            f"Loaded target_index {target_index} is out of bounds for y_train."
        )

    TARGET_CLASS = int(y_train[target_index])
    if TARGET_CLASS == 2:
        PERTURBING_CLASS = 1
        MISCLASSIFY_AS_CLASS = 1
    elif TARGET_CLASS == 0:
        PERTURBING_CLASS = 1
        MISCLASSIFY_AS_CLASS = 1
    elif TARGET_CLASS == 1:
        PERTURBING_CLASS = 0
        MISCLASSIFY_AS_CLASS = 0
    else:
        raise ValueError(
            f"Unexpected TARGET_CLASS {TARGET_CLASS} derived from index {target_index}"
        )

    print(f"Target Point True Class (TARGET_CLASS): {TARGET_CLASS}")
    print(f"Inferred Perturbing Class (PERTURBING_CLASS): {PERTURBING_CLASS}")
    print(
        f"Required Misclassification Class (MISCLASSIFY_AS_CLASS): {MISCLASSIFY_AS_CLASS}"
    )
    X_target_point = X_train[target_index]

except FileNotFoundError:
    print(f"{malware_red}Error:{white} Dataset file '{dataset_filename}' not found.")
    print("Make sure the .npz data file is in the correct directory.")
    raise
except KeyError as e:
    print(
        f"{malware_red}Error:{white} Could not find expected key '{e}' in '{dataset_filename}'."
    )
    raise
except Exception as e:
    print(f"{malware_red}An unexpected error occurred during data loading:{white} {e}")
    raise
Data loaded successfully from .npz file.
X_train shape: (1260, 2)
y_train shape: (1260,)
X_test shape: (540, 2)
y_test shape: (540,)
Target Point Index: 334
Target Point True Class (TARGET_CLASS): 2
Inferred Perturbing Class (PERTURBING_CLASS): 1
Required Misclassification Class (MISCLASSIFY_AS_CLASS): 1
In [3]:
def plot_data_multi(
    X,
    y,
    title="Multi-Class Dataset Visualization",
    highlight_indices=None,
    highlight_markers=None,
    highlight_colors=None,
    highlight_labels=None,
):
    """Plots 2D multi-class data with optional highlighting."""
    plt.figure(figsize=(12, 7))
    class_colors = [azure, nugget_yellow, malware_red]
    unique_classes = np.unique(y)
    max_class_idx = int(np.max(unique_classes)) if len(unique_classes) > 0 else -1
    if max_class_idx >= len(class_colors):
        class_colors.extend([hacker_grey] * (max_class_idx + 1 - len(class_colors)))
    cmap_multi = plt.cm.colors.ListedColormap(class_colors)

    plt.scatter(
        X[:, 0],
        X[:, 1],
        c=y,
        cmap=cmap_multi,
        edgecolors=node_black,
        s=50,
        alpha=0.7,
        zorder=1,
    )

    highlight_handles = []
    if highlight_indices is not None and len(highlight_indices) > 0:
        num_highlights = len(highlight_indices)
        _markers = highlight_markers if highlight_markers else ["o"] * num_highlights
        _colors = (
            highlight_colors if highlight_colors else [vivid_purple] * num_highlights
        )
        _labels = highlight_labels if highlight_labels else [""] * num_highlights
        for i, idx in enumerate(highlight_indices):
            if not (0 <= idx < X.shape[0]):
                continue
            marker = _markers[i % len(_markers)]
            edge_color = _colors[i % len(_colors)]
            label = _labels[i % len(_labels)]
            point_class = int(y[idx])
            face_color = (
                class_colors[point_class]
                if 0 <= point_class < len(class_colors)
                else hacker_grey
            )
            z_order = 3 if marker == "P" else 2
            plt.scatter(
                X[idx, 0],
                X[idx, 1],
                facecolors=face_color,
                edgecolors=edge_color,
                marker=marker,
                s=180,
                linewidths=2,
                alpha=1.0,
                zorder=z_order,
            )
            if label:
                highlight_handles.append(
                    plt.Line2D(
                        [0],
                        [0],
                        marker=marker,
                        color="w",
                        label=label,
                        markerfacecolor=face_color,
                        markeredgecolor=edge_color,
                        markersize=10,
                        linestyle="None",
                        markeredgewidth=1.5,
                    )
                )

    plt.title(title, fontsize=16, color=htb_green)
    plt.xlabel("Feature 1 (Standardized)", fontsize=12)
    plt.ylabel("Feature 2 (Standardized)", fontsize=12)
    class_handles = []
    unique_classes_present = sorted(np.unique(y))
    for class_idx in unique_classes_present:
        int_class_idx = int(class_idx)
        if 0 <= int_class_idx < len(class_colors):
            class_handles.append(
                plt.Line2D(
                    [0],
                    [0],
                    marker="o",
                    color="w",
                    label=f"Class {int_class_idx}",
                    markersize=10,
                    markerfacecolor=class_colors[int_class_idx],
                    markeredgecolor=node_black,
                    linestyle="None",
                )
            )
    all_handles = class_handles + highlight_handles
    if all_handles:
        plt.legend(handles=all_handles, title="Classes & Points")
    plt.grid(True, color=hacker_grey, linestyle="--", linewidth=0.5, alpha=0.3)
    plt.show()


def plot_decision_boundary_multi(
    X,
    y,
    model,
    title="Decision Boundary",
    highlight_indices=None,
    highlight_markers=None,
    highlight_colors=None,
    highlight_labels=None,
):
    """Plots decision boundaries for a multi-class model."""
    plt.figure(figsize=(12, 7))
    class_colors = [azure, nugget_yellow, malware_red]
    light_colors = [c + "60" for c in class_colors]  # Add alpha for contour fill
    max_class_data = int(np.max(y)) if len(y) > 0 else -1
    if max_class_data >= len(class_colors):
        class_colors.extend([hacker_grey] * (max_class_data + 1 - len(class_colors)))
        light_colors.extend(
            [hacker_grey + "60"] * (max_class_data + 1 - len(light_colors))
        )

    h = 0.02  # Mesh step size
    x_min, x_max = X[:, 0].min() - 0.5, X[:, 0].max() + 0.5
    y_min, y_max = X[:, 1].min() - 0.5, X[:, 1].max() + 0.5
    xx, yy = np.meshgrid(np.arange(x_min, x_max, h), np.arange(y_min, y_max, h))
    mesh_points = np.c_[xx.ravel(), yy.ravel()]

    try:
        Z = model.predict(mesh_points)
        Z = Z.reshape(xx.shape)
        max_class_pred = int(np.max(Z)) if Z.size > 0 else -1
        if max_class_pred >= len(class_colors):
            # Need to ensure enough colors for predicted classes on mesh
            needed = max_class_pred + 1
            if needed > len(class_colors):
                class_colors.extend([hacker_grey] * (needed - len(class_colors)))
                light_colors.extend([hacker_grey + "60"] * (needed - len(light_colors)))

        cmap_light = plt.cm.colors.ListedColormap(light_colors[: max_class_pred + 1])
        cmap_bold = plt.cm.colors.ListedColormap(class_colors[: max_class_data + 1])

        plt.contourf(xx, yy, Z, cmap=cmap_light, alpha=0.6, zorder=0)
        plt.scatter(
            X[:, 0],
            X[:, 1],
            c=y,
            cmap=cmap_bold,
            edgecolors=node_black,
            s=50,
            alpha=0.8,
            zorder=1,
        )

        # Plot highlighted points
        highlight_handles = []
        if highlight_indices is not None and len(highlight_indices) > 0:
            num_highlights = len(highlight_indices)
            _markers = (
                highlight_markers if highlight_markers else ["o"] * num_highlights
            )
            _colors = (
                highlight_colors
                if highlight_colors
                else [vivid_purple] * num_highlights
            )
            _labels = highlight_labels if highlight_labels else [""] * num_highlights
            for i, idx in enumerate(highlight_indices):
                if not (0 <= idx < X.shape[0]):
                    continue
                marker = _markers[i % len(_markers)]
                edge_color = _colors[i % len(_colors)]
                label = _labels[i % len(_labels)]
                point_class = int(y[idx])
                face_color = (
                    class_colors[point_class]
                    if 0 <= point_class < len(class_colors)
                    else hacker_grey
                )
                z_order = 3 if marker == "P" else 2
                plt.scatter(
                    X[idx, 0],
                    X[idx, 1],
                    facecolors=face_color,
                    edgecolors=edge_color,
                    marker=marker,
                    s=180,
                    linewidths=2,
                    alpha=1.0,
                    zorder=z_order,
                )
                if label:
                    highlight_handles.append(
                        plt.Line2D(
                            [0],
                            [0],
                            marker=marker,
                            color="w",
                            label=label,
                            markerfacecolor=face_color,
                            markeredgecolor=edge_color,
                            markersize=10,
                            linestyle="None",
                            markeredgewidth=1.5,
                        )
                    )

        plt.title(title, fontsize=16, color=htb_green)
        plt.xlabel("Feature 1 (Standardized)", fontsize=12)
        plt.ylabel("Feature 2 (Standardized)", fontsize=12)
        class_handles = []
        unique_classes_present = sorted(np.unique(y))
        for class_idx in unique_classes_present:
            int_class_idx = int(class_idx)
            if 0 <= int_class_idx < len(class_colors):
                class_handles.append(
                    plt.Line2D(
                        [0],
                        [0],
                        marker="o",
                        color="w",
                        label=f"Class {int_class_idx}",
                        markersize=10,
                        markerfacecolor=class_colors[int_class_idx],
                        markeredgecolor=node_black,
                        linestyle="None",
                    )
                )
        all_handles = class_handles + highlight_handles
        if all_handles:
            plt.legend(handles=all_handles, title="Classes & Points")

        plt.grid(True, color=hacker_grey, linestyle="--", linewidth=0.5, alpha=0.3)
        plt.xlim(xx.min(), xx.max())
        plt.ylim(yy.min(), yy.max())
        plt.show()

    except Exception as e:
        print(f"{malware_red}Error during plotting decision boundary:{white} {e}")
        plt.figure(figsize=(12, 7))
        plt.scatter(
            X[:, 0],
            X[:, 1],
            c=y,
            cmap=plt.cm.colors.ListedColormap(class_colors),
            edgecolors=node_black,
        )
        plt.title(f"{title} (Plotting Error)", color=malware_red)
        plt.show()


print("Visualization functions defined.")
Visualization functions defined.
In [4]:
print("\n--- Visualizing Clean Training Data with Target Point ---")
plot_data_multi(
    X_train,
    y_train,
    title=f"Clean Training Data (Target Idx: {target_index}, Class: {TARGET_CLASS})",
    highlight_indices=[target_index],
    highlight_markers=["P"],  # Plus sign for target
    highlight_colors=[white],  # White edge for visibility
    highlight_labels=[f"Target (Class {TARGET_CLASS}, Idx {target_index})"],
)
--- Visualizing Clean Training Data with Target Point ---
In [5]:
def perform_clean_label_attack(
    X_train_orig,
    y_train_orig,
    target_idx,
    target_class,
    perturb_class,
    n_neighbors,
    epsilon_cross,
    seed,
):
    # Implement the clean label attack logic here
    if target_class == perturb_class:
        raise ValueError("Target class and perturbing class cannot be the same.")

    # Create Copies to avoid modifying original data
    X_train_poisoned = X_train_orig.copy()
    y_train_poisoned = y_train_orig.copy() # Labels remain unchanged.

    # Train a temporary baseline model (OvR Logistic Regression)
    # We need this to find the decision boundary between target_class and perturb_class
    try:
        temp_base_estimator = LogisticRegression(
            random_state=seed, C=1.0, solver="liblinear"
        )
        temp_model = OneVsRestClassifier(temp_base_estimator)
        temp_model.fit(X_train_orig, y_train_orig)
        print("Temporary baseline model trained.")

        if not (
            hasattr(temp_model, "estimators_")
            and len(temp_model.estimators_) > max(target_class, perturb_class)
        ):
            raise RuntimeError(
                "Temporary model did not produce expected number of estimators."
            )

        w_target = temp_model.estimators_[target_class].coef_[0]
        b_target = temp_model.estimators_[target_class].intercept_[0]
        w_perturb = temp_model.estimators_[perturb_class].coef_[0]
        b_perturb = temp_model.estimators_[perturb_class].intercept_[0]
        print(
            f"Extracted weights/intercepts for Class {target_class} and Class {perturb_class}."
        )

    except Exception as e:
        print(f"Error: Failed to train or extract params from temporary model: {e}")
        raise RuntimeError("Failed to initialize temporary model for attack.") from e

    # Identify neighbors of perturb_class closest to the target point
    X_target_point = X_train_orig[target_idx]
    perturb_class_indices_train = np.where(y_train_orig == perturb_class)[0]

    if len(perturb_class_indices_train) == 0:
        raise ValueError(
            f"No points found for perturb_class ({perturb_class}). Cannot find neighbors."
        )
    if n_neighbors > len(perturb_class_indices_train):
        print(
            f"Warning: Requested {n_neighbors} neighbors, but only {len(perturb_class_indices_train)} points of Class {perturb_class} exist. Using all available."
        )
        n_neighbors = len(perturb_class_indices_train)
    if n_neighbors == 0:
        raise ValueError("n_neighbors is zero, cannot proceed with perturbation.")

    X_perturb_class_train = X_train_orig[perturb_class_indices_train]
    print(f"Finding {n_neighbors} nearest neighbors from Class {perturb_class}...")
    nn_finder = NearestNeighbors(n_neighbors=n_neighbors, algorithm="auto")
    nn_finder.fit(X_perturb_class_train)
    distances, indices_relative = nn_finder.kneighbors(X_target_point.reshape(1, -1))

    # Map relative indices back to original X_train indices
    neighbor_indices_absolute = perturb_class_indices_train[indices_relative.flatten()]
    X_neighbors_original = X_train_orig[neighbor_indices_absolute]
    print(f"Found neighbors at indices: {neighbor_indices_absolute}")

    # Calculate the perturbation vector
    #    Push direction is opposite to the normal vector of the boundary (w_target - w_perturb)
    w_diff_boundary = w_target - w_perturb
    b_diff_boundary = (
        b_target - b_perturb
    )  # Not needed for direction, but good for checks

    push_direction = w_diff_boundary
    norm_push_direction = np.linalg.norm(push_direction)

    if norm_push_direction < 1e-9:
        raise ValueError(
            "Boundary vector norm is close to zero. Cannot determine reliable push direction."
        )

    unit_push_direction = push_direction / norm_push_direction
    perturbation_vector = epsilon_cross * unit_push_direction
    print(f"Calculated perturbation vector (delta): {perturbation_vector}")

    # Apply perturbations to the neighbors' features in the copied dataset
    perturbed_indices_list = []
    print("Applying perturbations...")
    for i, neighbor_idx in enumerate(neighbor_indices_absolute):
        X_neighbor_orig = X_neighbors_original[i]
        X_perturbed_neighbor = X_neighbor_orig + perturbation_vector

        # Update the feature vector in the poisoned dataset
        X_train_poisoned[neighbor_idx] = X_perturbed_neighbor
        # DO NOT change y_train_poisoned[neighbor_idx]

        perturbed_indices_list.append(neighbor_idx)

        # Check if perturbation crossed the temporary baseline boundary
        f_boundary_orig = X_neighbor_orig @ w_diff_boundary + b_diff_boundary
        f_boundary_pert = X_perturbed_neighbor @ w_diff_boundary + b_diff_boundary
        # Expect f_boundary_orig < 0 (perturb class side), f_boundary_pert > 0 (target class side)
        print(
            f"  Neighbor {neighbor_idx}: Orig f={f_boundary_orig:.4f}, Perturbed f={f_boundary_pert:.4f}"
        )
        if f_boundary_pert <= 0:
            print(
                f"     Warning: Perturbed point {neighbor_idx} might not have crossed the boundary (f<=0)."
            )
            

    print(f"Applied perturbations to {len(perturbed_indices_list)} neighbors.")
    perturbed_indices = np.array(perturbed_indices_list)  # Return as numpy array

    # Final check: ensure target point wasn't accidentally perturbed
    if target_idx in perturbed_indices:
        print(
            f"CRITICAL Error: Target index {target_idx} was selected as a neighbor and perturbed! Check logic."
        )
        # Depending on desired strictness, could raise an error here

    print("--- Clean Label Attack Implementation Finished")
    
    return X_train_poisoned, y_train_poisoned, perturbed_indices
In [6]:
try:
    X_train_poisoned, y_train_poisoned, perturbed_indices = perform_clean_label_attack(
        X_train,
        y_train,
        target_index,
        TARGET_CLASS,
        PERTURBING_CLASS,
        N_NEIGHBORS,
        EPSILON_CROSS,
        SEED,
    )
    print(f"\nAttack function executed. Poisoned data shape: {X_train_poisoned.shape}")
    print(f"Indices perturbed: {perturbed_indices}")

    print("\n--- Training Final Model on Poisoned Data ---")
    poisoned_model = OneVsRestClassifier(
        LogisticRegression(random_state=SEED, C=1.0, solver="liblinear")
    )
    poisoned_model.fit(X_train_poisoned, y_train_poisoned)
    print("Final model trained successfully on poisoned data.")

    attack_logic_successful = True
except Exception as e:
    print(
        f"\n{malware_red}Error during attack execution or poisoned model training:{white} {e}"
    )
    print("Cannot proceed to evaluation or submission.")
    attack_logic_successful = False
    # Assign placeholder values to allow subsequent cells to run without crashing immediately
    X_train_poisoned, y_train_poisoned, perturbed_indices = (
        X_train,
        y_train,
        np.array([]),
    )
    poisoned_model = None  # Indicate model training failed
Temporary baseline model trained.
Extracted weights/intercepts for Class 2 and Class 1.
Finding 12 nearest neighbors from Class 1...
Found neighbors at indices: [1123  586 1214  555  596  785 1082 1231 1156  982 1122  385]
Calculated perturbation vector (delta): [0.10429172 0.22720747]
Applying perturbations...
  Neighbor 1123: Orig f=0.0098, Perturbed f=2.8168
  Neighbor 586: Orig f=-2.4651, Perturbed f=0.3420
  Neighbor 1214: Orig f=-2.6259, Perturbed f=0.1812
  Neighbor 555: Orig f=-2.3443, Perturbed f=0.4627
  Neighbor 596: Orig f=-2.6909, Perturbed f=0.1161
  Neighbor 785: Orig f=-2.6012, Perturbed f=0.2059
  Neighbor 1082: Orig f=-0.4691, Perturbed f=2.3379
  Neighbor 1231: Orig f=-2.9484, Perturbed f=-0.1414
     Warning: Perturbed point 1231 might not have crossed the boundary (f<=0).
  Neighbor 1156: Orig f=-2.6714, Perturbed f=0.1356
  Neighbor 982: Orig f=-3.3073, Perturbed f=-0.5003
     Warning: Perturbed point 982 might not have crossed the boundary (f<=0).
  Neighbor 1122: Orig f=2.8723, Perturbed f=5.6793
  Neighbor 385: Orig f=-2.9996, Perturbed f=-0.1925
     Warning: Perturbed point 385 might not have crossed the boundary (f<=0).
Applied perturbations to 12 neighbors.
--- Clean Label Attack Implementation Finished

Attack function executed. Poisoned data shape: (1260, 2)
Indices perturbed: [1123  586 1214  555  596  785 1082 1231 1156  982 1122  385]

--- Training Final Model on Poisoned Data ---
Final model trained successfully on poisoned data.
In [7]:
if attack_logic_successful and poisoned_model is not None:
    print("\n--- Evaluating Poisoned Model ---")

    # Check prediction for the target point
    X_target_reshaped = X_target_point.reshape(1, -1)
    target_pred_poisoned = poisoned_model.predict(X_target_reshaped)[0]

    print(f"Target Point Evaluation (Index: {target_index}):")
    print(f"  Original True Label:      {TARGET_CLASS}")
    print(f"  Poisoned Model Prediction: {target_pred_poisoned}")

    # Check if the misclassification matches the requirement
    attack_successful = target_pred_poisoned == MISCLASSIFY_AS_CLASS

    if attack_successful:
        print(
            f"  {htb_green}Success:{white} The poisoned model misclassified the target point as the required Class {MISCLASSIFY_AS_CLASS}."
        )
    else:
        if target_pred_poisoned == TARGET_CLASS:
            print(
                f"  {malware_red}Failure:{white} The poisoned model still correctly classified the target point as Class {target_pred_poisoned}."
            )
        else:
            print(
                f"  {nugget_yellow}Partial/Unexpected:{white} The poisoned model misclassified the target point as Class {target_pred_poisoned}, but NOT the required Class {MISCLASSIFY_AS_CLASS}."
            )

    # Evaluate overall accuracy on the clean test set (optional insight)
    try:
        y_pred_poisoned_test = poisoned_model.predict(X_test)
        poisoned_accuracy_test = accuracy_score(y_test, y_pred_poisoned_test)
        print(f"\nOverall Performance on Clean Test Set (for info):")
        print(f"  Poisoned Model Accuracy: {poisoned_accuracy_test:.4f}")
        print("\nClassification Report (Poisoned Model on Clean Test Data):")
        print(
            classification_report(
                y_test,
                y_pred_poisoned_test,
                target_names=[f"Class {i}" for i in range(3)],
            )
        )
    except Exception as e:
        print(
            f"\n{nugget_yellow}Warning:{white} Could not evaluate test set accuracy: {e}"
        )

    # --- Visualize Poisoned Data and Boundaries ---
    print("\n--- Visualizing Poisoned Training Data ---")
    # Highlight target and the points that were actually perturbed
    plot_data_multi(
        X_train_poisoned,
        y_train_poisoned,  # Labels are unchanged
        title=f"Poisoned Training Data (Perturbed Indices: {len(perturbed_indices)})",
        highlight_indices=[target_index] + perturbed_indices.tolist(),
        highlight_markers=["P"] + ["o"] * len(perturbed_indices),
        highlight_colors=[white]
        + [aquamarine] * len(perturbed_indices),  # Aquamarine for perturbed points
        highlight_labels=[f"Target (Idx {target_index})"]
        + [
            f"Perturbed (Idx {idx}, Label {PERTURBING_CLASS})"
            for idx in perturbed_indices
        ],
    )

    print("\n--- Visualizing Poisoned Model Decision Boundaries ---")
    plot_decision_boundary_multi(
        X_train_poisoned,  # Show poisoned points
        y_train_poisoned,
        poisoned_model,  # Use the poisoned model for boundaries
        title=f"Poisoned Model Decision Boundary\nTarget Pred: {target_pred_poisoned} | Required: {MISCLASSIFY_AS_CLASS}",
        highlight_indices=[target_index] + perturbed_indices.tolist(),
        highlight_markers=["P"] + ["o"] * len(perturbed_indices),
        highlight_colors=[white] + [aquamarine] * len(perturbed_indices),
        highlight_labels=[f"Target (Pred: {target_pred_poisoned})"]
        + [f"Perturbed (Idx {idx})" for idx in perturbed_indices],
    )

else:
    print(
        "\nSkipping evaluation and visualization due to errors in attack implementation or model training."
    )
--- Evaluating Poisoned Model ---
Target Point Evaluation (Index: 334):
  Original True Label:      2
  Poisoned Model Prediction: 1
  #9fef00Success:#ffffff The poisoned model misclassified the target point as the required Class 1.

Overall Performance on Clean Test Set (for info):
  Poisoned Model Accuracy: 0.9815

Classification Report (Poisoned Model on Clean Test Data):
              precision    recall  f1-score   support

     Class 0       1.00      1.00      1.00       180
     Class 1       0.97      0.98      0.97       180
     Class 2       0.98      0.97      0.97       180

    accuracy                           0.98       540
   macro avg       0.98      0.98      0.98       540
weighted avg       0.98      0.98      0.98       540


--- Visualizing Poisoned Training Data ---
--- Visualizing Poisoned Model Decision Boundaries ---
In [8]:
# Extract weights and intercept from the successfully trained poisoned model
if attack_logic_successful and poisoned_model is not None:
    try:
        # For OvR, we need weights/intercepts from all base estimators
        # The evaluator expects a list of lists for weights and a list for intercepts
        weights_list = [est.coef_[0].tolist() for est in poisoned_model.estimators_]
        intercept_list = [est.intercept_[0] for est in poisoned_model.estimators_]
        print("Extracted weights and intercepts from the poisoned model.")
        print(
            f"Weights shape (list of lists): ({len(weights_list)}, {len(weights_list[0]) if weights_list else 0})"
        )
        print(f"Intercept shape (list): ({len(intercept_list)})")
        submission_ready = True
    except Exception as e:
        print(
            f"{malware_red}Error:{white} Failed to extract parameters from poisoned model: {e}"
        )
        submission_ready = False
else:
    print(
        "Skipping parameter extraction because the model was not trained successfully."
    )
    submission_ready = False
Extracted weights and intercepts from the poisoned model.
Weights shape (list of lists): (3, 2)
Intercept shape (list): (3)
In [9]:
# Replace <EVALUATOR_IP> and <PORT> with the correct values for the lab environment.
# Example: evaluator_base_url = "http://10.10.10.5:5000"
# evaluator_base_url = "http://<EVALUATOR_IP>:<PORT>"
evaluator_base_url = "http://154.57.164.62:30910"

# --- Health Check ---
health_check_url = f"{evaluator_base_url}/health"
print(f"Checking evaluator health at: {health_check_url}")

if "<EVALUATOR_IP>" in evaluator_base_url:
    print(f"\n{nugget_yellow}--- WARNING ---")
    print(
        "Please update the 'evaluator_base_url' variable with the correct IP and Port before running!"
    )
    print("----------------{white}")
else:
    try:
        response = requests.get(health_check_url, timeout=10)
        response.raise_for_status()  # Raise HTTPError for bad responses (4xx or 5xx)
        health_status = response.json()
        print("\n--- Health Check Response ---")
        print(f"Status: {health_status.get('status', 'N/A')}")
        print(f"Message: {health_status.get('message', 'No message received.')}")
        if health_status.get("status") != "healthy":
            print(
                f"\n{nugget_yellow}Warning:{white} Evaluator service reported an unhealthy status. Submission might fail."
            )
    except requests.exceptions.ConnectionError:
        print(
            f"\n{malware_red}Connection Error:{white} Could not connect to {health_check_url}."
        )
        print(
            "Please check the evaluator URL (IP/Port) and ensure the Docker container is running."
        )
    except requests.exceptions.Timeout:
        print(
            f"\n{malware_red}Timeout Error:{white} The request to {health_check_url} timed out."
        )
    except requests.exceptions.RequestException as e:
        print(f"\n{malware_red}Error during health check:{white} {e}")
    except json.JSONDecodeError:
        print(
            f"\n{malware_red}Error:{white} Could not decode JSON response from health check."
        )
        print(f"Raw status: {response.status_code}, Raw text: {response.text}")
    except Exception as e:
        print(
            f"\n{malware_red}An unexpected error occurred during health check:{white} {e}"
        )

# --- Submission ---
evaluator_url = f"{evaluator_base_url}/evaluate"
print(f"\nAttempting submission to: {evaluator_url}")

if "<EVALUATOR_IP>" in evaluator_base_url:
    print(f"\n{nugget_yellow}--- WARNING ---")
    print("Update 'evaluator_base_url' with the correct IP and Port before submitting!")
    print("----------------{white}")
elif not submission_ready:
    print(
        f"\n{malware_red}Submission aborted:{white} Model parameters were not extracted successfully."
    )
else:
    # Prepare payload
    payload = {"weights": weights_list, "intercept": intercept_list}
    # print(f"Payload preview: {json.dumps(payload, indent=2)[:500]}...")

    try:
        response = requests.post(evaluator_url, json=payload, timeout=30)
        response.raise_for_status()  # Raise HTTPError for bad responses (4xx or 5xx)
        result = response.json()

        print("\n--- Evaluator Response ---")
        if result.get("success"):
            print(f"{htb_green}Attack Successful!{white}")
            # Accuracy reported by server might be on the *poisoned* test set or a hidden one
            print(f"Accuracy reported by server: {result.get('accuracy', 'N/A'):.4f}")
            print(f"Target Point Check: {result.get('message', 'No message.')}")
            print(f"{htb_green}Flag: {result.get('flag', 'FLAG_NOT_RETURNED')}{white}")
        else:
            print(f"{malware_red}Evaluation Failed.{white}")
            accuracy_val = result.get("accuracy")
            accuracy_str = f"{accuracy_val:.4f}" if accuracy_val is not None else "N/A"
            print(f"Accuracy reported by server: {accuracy_str}")
            print(f"Message: {result.get('message', 'No specific message.')}")
            print("\nHints: Did your attack successfully misclassify the target point?")
            print(
                "Did you perturb the correct neighbors? Is epsilon_cross appropriate?"
            )
            print("Did you submit parameters for a 3-class OvR model?")

    except requests.exceptions.ConnectionError:
        print(
            f"\n{malware_red}Connection Error:{white} Could not connect to {evaluator_url}."
        )
        print("Check the evaluator URL and ensure the Docker container is running.")
    except requests.exceptions.Timeout:
        print(
            f"\n{malware_red}Timeout Error:{white} The request to {evaluator_url} timed out."
        )
    except requests.exceptions.RequestException as e:
        # This catches HTTPError (4xx, 5xx) as well
        print(f"\n{malware_red}Error during submission:{white} {e}")
        try:
            # Try to get more info from the response if available
            error_details = response.json()
            print(f"Server responded with: {error_details}")
        except (AttributeError, json.JSONDecodeError):
            if hasattr(response, "text"):
                print(f"Raw response content: {response.text}")
            else:
                print("(No response content available)")
    except json.JSONDecodeError:
        print(
            f"\n{malware_red}Error:{white} Could not decode JSON response from evaluator."
        )
        print(f"Raw status: {response.status_code}, Raw text: {response.text}")
    except Exception as e:
        print(
            f"\n{malware_red}An unexpected error occurred during submission:{white} {e}"
        )
Checking evaluator health at: http://154.57.164.62:30910/health

--- Health Check Response ---
Status: healthy
Message: Evaluator API is running.

Attempting submission to: http://154.57.164.62:30910/evaluate

--- Evaluator Response ---
#9fef00Attack Successful!#ffffff
Accuracy reported by server: 0.9815
Target Point Check: Attack successful! Target point (Index 334, True Class 2) misclassified as 1. Overall accuracy (0.9815) maintained and sensitivity checks passed.
#9fef00Flag: HTB{cl3an_l4b3l_fl4g_fun}#ffffff