Files
AI-Red-Teaming-CSCD94/defense/adversarial training/adversarial training.ipynb
T
2026-07-26 22:53:03 -04:00

30 KiB

In [1]:
import json
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from safetensors.torch import save_file
from tqdm import tqdm

from htb_ai_library import (
    set_reproducibility,
    get_mnist_loaders,
    evaluate_accuracy,
    train_model,
)

MNIST_MEAN = 0.1307
MNIST_STD = 0.3081
EPSILON = 0.3
EPSILON_SPREAD = [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0]
I_FGSM_STEPS = 10


def get_device():
    """Get the best available device."""
    if torch.cuda.is_available():
        return torch.device("cuda")
    return torch.device("cpu")

def save_adversarial_examples(data, path):
    """Save adversarial examples to safetensors format."""
    # We don't store fgsm_images/ifgsm_images separately since they reference
    # fgsm_by_epsilon[epsilon] and would cause memory sharing errors in safetensors
    tensors = {
        'clean_images': data['clean_images'],
        'clean_labels': data['clean_labels'].long(),
    }

    for eps in data['epsilon_spread']:
        key = f"{eps:.1f}"
        tensors[f'fgsm_eps_{key}'] = data['fgsm_by_epsilon'][eps]
        tensors[f'ifgsm_eps_{key}'] = data['ifgsm_by_epsilon'][eps]

    metadata = {
        'epsilon': str(data['epsilon']),
        'epsilon_spread': json.dumps(data['epsilon_spread']),
    }

    save_file(tensors, path, metadata=metadata)

class LeNet5(nn.Module):
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 6, kernel_size=5, padding=2)
        self.conv2 = nn.Conv2d(6, 16, kernel_size=5)
        self.fc1 = nn.Linear(16 * 5 * 5, 120)
        self.fc2 = nn.Linear(120, 84)
        self.fc3 = nn.Linear(84, 10)

    def forward(self, x):
        x = F.max_pool2d(F.relu(self.conv1(x)), 2)
        x = F.max_pool2d(F.relu(self.conv2(x)), 2)
        x = x.view(-1, 16 * 5 * 5)
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        x = self.fc3(x)
        return x

def fgsm_attack(model, images, labels, epsilon):
    images_copy = images.clone().detach().requires_grad_(True)
    outputs = model(images_copy)
    loss = F.cross_entropy(outputs, labels)
    model.zero_grad()
    loss.backward()
    grad_sign = images_copy.grad.sign()
    adv_images = images_copy + epsilon * grad_sign
    min_val = (0 - MNIST_MEAN) / MNIST_STD
    max_val = (1 - MNIST_MEAN) / MNIST_STD
    adv_images = torch.clamp(adv_images, min_val, max_val)
    return adv_images.detach()

def i_fgsm_attack(model, images, labels, epsilon, steps=I_FGSM_STEPS):
    """Generate I-FGSM (Iterative FGSM) adversarial examples."""
    alpha = epsilon / steps  # Step size per iteration
    min_val = (0 - MNIST_MEAN) / MNIST_STD
    max_val = (1 - MNIST_MEAN) / MNIST_STD

    adv_images = images.clone().detach()
    original_images = images.clone().detach()

    for _ in range(steps):
        adv_images.requires_grad = True
        outputs = model(adv_images)
        loss = F.cross_entropy(outputs, labels)
        model.zero_grad()
        loss.backward()

        grad_sign = adv_images.grad.sign()
        adv_images = adv_images.detach() + alpha * grad_sign

        # Project back to epsilon-ball around original
        perturbation = adv_images - original_images
        perturbation = torch.clamp(perturbation, -epsilon, epsilon)
        adv_images = original_images + perturbation

        # Clamp to valid range
        adv_images = torch.clamp(adv_images, min_val, max_val)

    return adv_images.detach()

def evaluate_adversarial_accuracy(model, loader, device, epsilon, num_batches=None):
    """Evaluate accuracy under FGSM attack."""
    model.eval()
    correct = 0
    total = 0

    for i, (images, labels) in enumerate(loader):
        if num_batches is not None and i >= num_batches:
            break

        images, labels = images.to(device), labels.to(device)

        # Generate adversarial examples (need gradients, so briefly enable train mode)
        model.train()
        adv_images = fgsm_attack(model, images, labels, epsilon)
        model.eval()

        with torch.no_grad():
            outputs = model(adv_images)
            _, predicted = outputs.max(1)
            total += labels.size(0)
            correct += predicted.eq(labels).sum().item()

    return 100.0 * correct / total

train_loader, test_loader = get_mnist_loaders(batch_size=128, data_dir="./data")
baseline_model = LeNet5()
baseline_model = train_model(
    baseline_model,
    train_loader,
    test_loader,
    device=get_device(),
    epochs=10,
    learning_rate=0.001,
)
Downloading http://yann.lecun.com/exdb/mnist/train-images-idx3-ubyte.gz
Failed to download (trying next):
HTTP Error 404: Not Found

Downloading https://ossci-datasets.s3.amazonaws.com/mnist/train-images-idx3-ubyte.gz
Downloading https://ossci-datasets.s3.amazonaws.com/mnist/train-images-idx3-ubyte.gz to ./data/MNIST/raw/train-images-idx3-ubyte.gz
100%|████████████████████████████████████████████████████████████████████████████████████████████████████████| 9.91M/9.91M [00:00<00:00, 51.7MB/s]
Extracting ./data/MNIST/raw/train-images-idx3-ubyte.gz to ./data/MNIST/raw

Downloading http://yann.lecun.com/exdb/mnist/train-labels-idx1-ubyte.gz
Failed to download (trying next):
HTTP Error 404: Not Found

Downloading https://ossci-datasets.s3.amazonaws.com/mnist/train-labels-idx1-ubyte.gz
Downloading https://ossci-datasets.s3.amazonaws.com/mnist/train-labels-idx1-ubyte.gz to ./data/MNIST/raw/train-labels-idx1-ubyte.gz
100%|████████████████████████████████████████████████████████████████████████████████████████████████████████| 28.9k/28.9k [00:00<00:00, 1.38MB/s]
Extracting ./data/MNIST/raw/train-labels-idx1-ubyte.gz to ./data/MNIST/raw

Downloading http://yann.lecun.com/exdb/mnist/t10k-images-idx3-ubyte.gz
Failed to download (trying next):
HTTP Error 404: Not Found

Downloading https://ossci-datasets.s3.amazonaws.com/mnist/t10k-images-idx3-ubyte.gz
Downloading https://ossci-datasets.s3.amazonaws.com/mnist/t10k-images-idx3-ubyte.gz to ./data/MNIST/raw/t10k-images-idx3-ubyte.gz
100%|████████████████████████████████████████████████████████████████████████████████████████████████████████| 1.65M/1.65M [00:00<00:00, 11.4MB/s]
Extracting ./data/MNIST/raw/t10k-images-idx3-ubyte.gz to ./data/MNIST/raw

Downloading http://yann.lecun.com/exdb/mnist/t10k-labels-idx1-ubyte.gz
Failed to download (trying next):
HTTP Error 404: Not Found

Downloading https://ossci-datasets.s3.amazonaws.com/mnist/t10k-labels-idx1-ubyte.gz
Downloading https://ossci-datasets.s3.amazonaws.com/mnist/t10k-labels-idx1-ubyte.gz to ./data/MNIST/raw/t10k-labels-idx1-ubyte.gz
100%|████████████████████████████████████████████████████████████████████████████████████████████████████████| 4.54k/4.54k [00:00<00:00, 5.35MB/s]
Extracting ./data/MNIST/raw/t10k-labels-idx1-ubyte.gz to ./data/MNIST/raw

Epoch 1/10: Avg Loss = 0.3898, Test Accuracy = 96.67%
Epoch 2/10: Avg Loss = 0.1010, Test Accuracy = 97.92%
Epoch 3/10: Avg Loss = 0.0693, Test Accuracy = 98.43%
Epoch 4/10: Avg Loss = 0.0523, Test Accuracy = 98.18%
Epoch 5/10: Avg Loss = 0.0448, Test Accuracy = 98.82%
Epoch 6/10: Avg Loss = 0.0368, Test Accuracy = 98.56%
Epoch 7/10: Avg Loss = 0.0327, Test Accuracy = 98.95%
Epoch 8/10: Avg Loss = 0.0280, Test Accuracy = 99.02%
Epoch 9/10: Avg Loss = 0.0254, Test Accuracy = 98.89%
Epoch 10/10: Avg Loss = 0.0211, Test Accuracy = 99.02%
In [2]:
def generate_adversarial_examples(model, test_loader, device, num_samples=500):
    """Generate adversarial examples across multiple epsilon values."""
    model.eval()

    # Collect clean samples
    clean_images_list = []
    clean_labels_list = []
    collected = 0

    print(f"Collecting {num_samples} clean samples...")
    for images, labels in test_loader:
        if collected >= num_samples:
            break
        batch_size = min(images.size(0), num_samples - collected)
        clean_images_list.append(images[:batch_size])
        clean_labels_list.append(labels[:batch_size])
        collected += batch_size

    clean_images = torch.cat(clean_images_list, dim=0)
    clean_labels = torch.cat(clean_labels_list, dim=0)

    # Generate adversarial examples at each epsilon
    fgsm_by_epsilon = {}
    ifgsm_by_epsilon = {}

    print(f"\nGenerating adversarial examples across epsilon spread: {EPSILON_SPREAD}")

    for eps in EPSILON_SPREAD:
        print(f"\n  Generating at epsilon={eps}...")
        fgsm_images_list = []
        ifgsm_images_list = []

        batch_size = 128
        pbar = tqdm(total=num_samples, desc=f"  eps={eps}")

        for i in range(0, num_samples, batch_size):
            end_idx = min(i + batch_size, num_samples)
            images = clean_images[i:end_idx].to(device)
            labels = clean_labels[i:end_idx].to(device)

            model.train()  # Need train mode for gradient computation
            fgsm_images = fgsm_attack(model, images, labels, eps)
            ifgsm_images = i_fgsm_attack(model, images, labels, eps)
            model.eval()

            fgsm_images_list.append(fgsm_images.cpu())
            ifgsm_images_list.append(ifgsm_images.cpu())
            pbar.update(end_idx - i)

        pbar.close()

        fgsm_by_epsilon[eps] = torch.cat(fgsm_images_list, dim=0)
        ifgsm_by_epsilon[eps] = torch.cat(ifgsm_images_list, dim=0)

    return {
        'clean_images': clean_images,
        'clean_labels': clean_labels,
        'fgsm_images': fgsm_by_epsilon[EPSILON],
        'ifgsm_images': ifgsm_by_epsilon[EPSILON],
        'epsilon': EPSILON,
        'epsilon_spread': EPSILON_SPREAD,
        'fgsm_by_epsilon': fgsm_by_epsilon,
        'ifgsm_by_epsilon': ifgsm_by_epsilon
    }

def train_adversarial(model, train_loader, test_loader, device,
                      epochs=25, lr=0.001, epsilon=EPSILON):
    """Train model with adversarial training using epsilon spread."""
    model.to(device)

    optimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)
    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)
    criterion = nn.CrossEntropyLoss()

    for epoch in range(epochs):
        model.train()
        total_loss = 0.0
        correct = 0
        total = 0

        pbar = tqdm(train_loader, desc=f"Epoch {epoch+1}/{epochs}")
        for images, labels in pbar:
            images, labels = images.to(device), labels.to(device)

            batch_epsilon = np.random.choice(EPSILON_SPREAD)
            adv_images = fgsm_attack(model, images, labels, batch_epsilon)

            combined_images = torch.cat([images, adv_images], dim=0)
            combined_labels = torch.cat([labels, labels], dim=0)

            perm = torch.randperm(combined_images.size(0))
            combined_images = combined_images[perm]
            combined_labels = combined_labels[perm]

            optimizer.zero_grad()
            outputs = model(combined_images)
            loss = criterion(outputs, combined_labels)
            loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
            optimizer.step()

            total_loss += loss.item()
            _, predicted = outputs.max(1)
            total += combined_labels.size(0)
            correct += predicted.eq(combined_labels).sum().item()

            pbar.set_postfix({
                'loss': f'{total_loss / (pbar.n + 1):.4f}',
                'acc': f'{100.0 * correct / total:.2f}%'
            })

        scheduler.step()

        if (epoch + 1) % 5 == 0 or epoch == epochs - 1:
            clean_acc = evaluate_accuracy(model, test_loader, device)
            adv_acc = evaluate_adversarial_accuracy(
                model, test_loader, device, epsilon, num_batches=20
            )
            print(f"\n  Epoch {epoch+1}: Clean={clean_acc:.1f}%, Robust={adv_acc:.1f}%")

    return model

set_reproducibility(1337)
device = get_device()
print(f"Using device: {device}")

train_loader, test_loader = get_mnist_loaders(normalize=True)
print(f"Train batches: {len(train_loader)}, Test batches: {len(test_loader)}")
Using device: cpu
Train batches: 469, Test batches: 79
In [3]:
baseline_model = LeNet5()
baseline_model = train_model(baseline_model, train_loader, test_loader,
                              device=device, epochs=10, learning_rate=0.001)

adv_acc = evaluate_adversarial_accuracy(baseline_model, test_loader, device, EPSILON)
print(f"Baseline Model - Robust: {adv_acc:.1f}%")
save_file(baseline_model.state_dict(), "baseline_model.safetensors")
Epoch 1/10: Avg Loss = 0.3098, Test Accuracy = 97.53%
Epoch 2/10: Avg Loss = 0.0809, Test Accuracy = 98.48%
Epoch 3/10: Avg Loss = 0.0539, Test Accuracy = 98.63%
Epoch 4/10: Avg Loss = 0.0431, Test Accuracy = 98.49%
Epoch 5/10: Avg Loss = 0.0348, Test Accuracy = 98.78%
Epoch 6/10: Avg Loss = 0.0295, Test Accuracy = 98.71%
Epoch 7/10: Avg Loss = 0.0257, Test Accuracy = 98.78%
Epoch 8/10: Avg Loss = 0.0204, Test Accuracy = 98.84%
Epoch 9/10: Avg Loss = 0.0167, Test Accuracy = 98.96%
Epoch 10/10: Avg Loss = 0.0151, Test Accuracy = 98.88%
Baseline Model - Robust: 74.8%
In [4]:
adv_data = generate_adversarial_examples(
    baseline_model, test_loader, device, num_samples=500
)

save_adversarial_examples(adv_data, "adv_examples.safetensors")
print("Adversarial examples saved to adv_examples.safetensors")

model = LeNet5()
model.to(device)

clean_acc = evaluate_accuracy(model, test_loader, device)
adv_acc = evaluate_adversarial_accuracy(model, test_loader, device, EPSILON)
print(f"Before training - Clean: {clean_acc:.1f}%, Robust: {adv_acc:.1f}%")

model = train_adversarial(model, train_loader, test_loader, device, epochs=10)

clean_acc = evaluate_accuracy(model, test_loader, device)
adv_acc = evaluate_adversarial_accuracy(model, test_loader, device, EPSILON)
print(f"Final - Clean: {clean_acc:.1f}%, Robust: {adv_acc:.1f}%")

save_file(model.state_dict(), "robust_model.safetensors")
print("Model saved to robust_model.safetensors")
Collecting 500 clean samples...

Generating adversarial examples across epsilon spread: [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0]

  Generating at epsilon=0.1...
  eps=0.1: 100%|██████████████████████████████████████████████████████████████████████████████████████████████| 500/500 [00:00<00:00, 3589.91it/s]
  Generating at epsilon=0.2...
  eps=0.2: 100%|██████████████████████████████████████████████████████████████████████████████████████████████| 500/500 [00:00<00:00, 3783.98it/s]
  Generating at epsilon=0.3...
  eps=0.3: 100%|██████████████████████████████████████████████████████████████████████████████████████████████| 500/500 [00:00<00:00, 4337.29it/s]
  Generating at epsilon=0.4...
  eps=0.4: 100%|██████████████████████████████████████████████████████████████████████████████████████████████| 500/500 [00:00<00:00, 4218.09it/s]
  Generating at epsilon=0.5...
  eps=0.5: 100%|██████████████████████████████████████████████████████████████████████████████████████████████| 500/500 [00:00<00:00, 2903.09it/s]
  Generating at epsilon=0.6...
  eps=0.6: 100%|██████████████████████████████████████████████████████████████████████████████████████████████| 500/500 [00:00<00:00, 2815.23it/s]
  Generating at epsilon=0.7...
  eps=0.7: 100%|██████████████████████████████████████████████████████████████████████████████████████████████| 500/500 [00:00<00:00, 2339.08it/s]
  Generating at epsilon=0.8...
  eps=0.8: 100%|██████████████████████████████████████████████████████████████████████████████████████████████| 500/500 [00:00<00:00, 2391.63it/s]
  Generating at epsilon=0.9...
  eps=0.9: 100%|██████████████████████████████████████████████████████████████████████████████████████████████| 500/500 [00:00<00:00, 2493.57it/s]
  Generating at epsilon=1.0...
  eps=1.0: 100%|██████████████████████████████████████████████████████████████████████████████████████████████| 500/500 [00:00<00:00, 2372.05it/s]
Adversarial examples saved to adv_examples.safetensors
Before training - Clean: 10.2%, Robust: 2.4%
Epoch 1/10: 100%|██████████████████████████████████████████████████████████████████████| 469/469 [00:06<00:00, 70.61it/s, loss=0.7535, acc=75.68%]
Epoch 2/10: 100%|██████████████████████████████████████████████████████████████████████| 469/469 [00:06<00:00, 73.95it/s, loss=0.3464, acc=88.43%]
Epoch 3/10: 100%|██████████████████████████████████████████████████████████████████████| 469/469 [00:06<00:00, 71.96it/s, loss=0.2837, acc=90.47%]
Epoch 4/10: 100%|██████████████████████████████████████████████████████████████████████| 469/469 [00:06<00:00, 72.84it/s, loss=0.2241, acc=92.43%]
Epoch 5/10: 100%|██████████████████████████████████████████████████████████████████████| 469/469 [00:06<00:00, 72.62it/s, loss=0.1920, acc=93.46%]
  Epoch 5: Clean=98.9%, Robust=92.9%
Epoch 6/10: 100%|██████████████████████████████████████████████████████████████████████| 469/469 [00:06<00:00, 71.49it/s, loss=0.1725, acc=94.28%]
Epoch 7/10: 100%|██████████████████████████████████████████████████████████████████████| 469/469 [00:06<00:00, 72.60it/s, loss=0.1551, acc=94.83%]
Epoch 8/10: 100%|██████████████████████████████████████████████████████████████████████| 469/469 [00:06<00:00, 73.79it/s, loss=0.1474, acc=95.11%]
Epoch 9/10: 100%|██████████████████████████████████████████████████████████████████████| 469/469 [00:06<00:00, 73.62it/s, loss=0.1393, acc=95.43%]
Epoch 10/10: 100%|█████████████████████████████████████████████████████████████████████| 469/469 [00:06<00:00, 72.08it/s, loss=0.1319, acc=95.60%]
  Epoch 10: Clean=99.0%, Robust=93.8%
Final - Clean: 99.0%, Robust: 95.7%
Model saved to robust_model.safetensors
In [ ]: