210 lines
6.6 KiB
Python
210 lines
6.6 KiB
Python
from __future__ import annotations
|
|
|
|
import argparse
|
|
import base64
|
|
import io
|
|
import json
|
|
from dataclasses import dataclass
|
|
import os
|
|
import time
|
|
from typing import Tuple
|
|
import numpy as np
|
|
import requests
|
|
from PIL import Image
|
|
import torch
|
|
import torch.nn as nn
|
|
|
|
# Define MNIST normalization constants
|
|
MNIST_MEAN = 0.1307 # average pixel intensity of MNIST images scaled to [0,1]
|
|
MNIST_STD = 0.3081 # standard deviation of pixel intensities in [0,1]
|
|
|
|
class SimpleClassifier(nn.Module):
|
|
"""CNN matching the server-side architecture with log-softmax outputs."""
|
|
|
|
def __init__(self) -> None:
|
|
super().__init__()
|
|
self.conv1 = nn.Conv2d(1, 32, 3, 1)
|
|
self.conv2 = nn.Conv2d(32, 64, 3, 1)
|
|
self.dropout1 = nn.Dropout(0.25)
|
|
self.dropout2 = nn.Dropout(0.5)
|
|
self.fc1 = nn.Linear(9216, 128)
|
|
self.fc2 = nn.Linear(128, 10)
|
|
|
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
|
x = self.conv1(x)
|
|
x = torch.relu(x)
|
|
x = self.conv2(x)
|
|
x = torch.relu(x)
|
|
x = torch.max_pool2d(x, 2)
|
|
x = self.dropout1(x)
|
|
x = torch.flatten(x, 1)
|
|
x = self.fc1(x)
|
|
x = torch.relu(x)
|
|
x = self.dropout2(x)
|
|
x = self.fc2(x)
|
|
return torch.log_softmax(x, dim=1)
|
|
|
|
|
|
def mnist_normalize(x01: torch.Tensor) -> torch.Tensor:
|
|
"""Normalize a [0,1] tensor to MNIST stats for the classifier."""
|
|
return (x01 - MNIST_MEAN) / MNIST_STD
|
|
|
|
def png_from_x01(x01: np.ndarray) -> str:
|
|
"""Encode a `[0,1]` grayscale image `(28,28)` to base64 PNG string."""
|
|
x255 = np.clip((x01 * 255.0).round(), 0, 255).astype(np.uint8)
|
|
img = Image.fromarray(x255, mode="L")
|
|
buf = io.BytesIO()
|
|
img.save(buf, format="PNG", optimize=True)
|
|
return base64.b64encode(buf.getvalue()).decode("ascii")
|
|
|
|
|
|
def x01_from_b64_png(b64: str) -> np.ndarray:
|
|
"""Decode base64 PNG to `[0,1]` numpy array of shape `(28,28)`."""
|
|
raw = base64.b64decode(b64)
|
|
img = Image.open(io.BytesIO(raw)).convert("L")
|
|
if img.size != (28, 28):
|
|
raise ValueError("Expected 28x28 PNG")
|
|
x = np.asarray(img, dtype=np.float32) / 255.0
|
|
return np.clip(x, 0.0, 1.0)
|
|
|
|
@dataclass
|
|
class Challenge:
|
|
epsilon: float
|
|
label: int
|
|
sample_index: int
|
|
x01: np.ndarray # (1,1,28,28)
|
|
|
|
def fetch_challenge(host: str, retries: int = 30, delay: float = 1.0) -> Challenge:
|
|
"""Fetch challenge with simple retry/backoff to tolerate startup races."""
|
|
last_err = None
|
|
for _ in range(max(1, retries)):
|
|
try:
|
|
r = requests.get(f"{host}/challenge", timeout=5)
|
|
r.raise_for_status()
|
|
payload = r.json()
|
|
break
|
|
except Exception as e: # noqa: BLE001
|
|
last_err = e
|
|
time.sleep(delay)
|
|
else:
|
|
raise RuntimeError(
|
|
f"Failed to connect to {host}/challenge: {last_err}"
|
|
) from last_err
|
|
x2d = x01_from_b64_png(payload["image_b64"]) # (28,28)
|
|
x4d = x2d[None, None, ...]
|
|
return Challenge(
|
|
epsilon=float(payload["epsilon"]),
|
|
label=int(payload["label"]),
|
|
sample_index=int(payload["sample_index"]),
|
|
x01=x4d.astype(np.float32),
|
|
)
|
|
|
|
def load_model(weights_path: str) -> SimpleClassifier:
|
|
model = SimpleClassifier()
|
|
state = torch.load(weights_path, map_location=torch.device("cpu"))
|
|
model.load_state_dict(state)
|
|
model.eval()
|
|
return model
|
|
|
|
def fgsm_untargeted(
|
|
model: nn.Module, x01: np.ndarray, y: int, epsilon: float
|
|
) -> np.ndarray:
|
|
"""Craft FGSM adversarial example under L_inf in [0,1].
|
|
|
|
Parameters
|
|
----------
|
|
model : nn.Module
|
|
Classifier in eval mode.
|
|
x01 : np.ndarray
|
|
Baseline `[0,1]` image with shape `(1,1,28,28)`.
|
|
y : int
|
|
True class label.
|
|
epsilon : float
|
|
L_inf bound in `[0,1]` space.
|
|
|
|
Returns
|
|
-------
|
|
np.ndarray
|
|
Adversarial image in `[0,1]` with shape `(1,1,28,28)`.
|
|
"""
|
|
x = torch.from_numpy(x01).float().requires_grad_(True)
|
|
target = torch.tensor([y], dtype=torch.long)
|
|
logits = model(mnist_normalize(x))
|
|
loss = nn.NLLLoss()(logits, target)
|
|
model.zero_grad()
|
|
loss.backward()
|
|
grad_sign = x.grad.detach().sign()
|
|
x_adv = x + epsilon * grad_sign
|
|
x_adv = torch.clamp(x_adv, 0.0, 1.0)
|
|
return x_adv.detach().cpu().numpy()
|
|
|
|
def main() -> None:
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument(
|
|
"--host", default="http://154.57.164.80:31923", help="Server base URL"
|
|
)
|
|
parser.add_argument(
|
|
"--weights", default="solver/fgsm_weights.pth", help="Path to model weights (downloads from /weights if missing)"
|
|
)
|
|
args = parser.parse_args()
|
|
|
|
chall = fetch_challenge(args.host)
|
|
if not os.path.exists(args.weights):
|
|
os.makedirs(os.path.dirname(args.weights), exist_ok=True)
|
|
wb = requests.get(f"{args.host}/weights", timeout=15).content
|
|
with open(args.weights, "wb") as f:
|
|
f.write(wb)
|
|
model = load_model(args.weights)
|
|
|
|
# Local sanity prediction on clean image
|
|
x = torch.from_numpy(chall.x01)
|
|
clean_pred = int(torch.argmax(model(mnist_normalize(x)), dim=1).item())
|
|
if clean_pred != chall.label:
|
|
print(
|
|
f"Warning: local clean prediction {clean_pred} != label {chall.label}; proceeding regardless"
|
|
)
|
|
|
|
# Try conservative epsilons to account for 8-bit PNG quantization
|
|
safety = [1.0 / 255.0 + 1e-4, 2e-3, 1e-3, 5e-4, 0.0]
|
|
x_adv = None
|
|
adv_pred = None
|
|
linf = None
|
|
for m in safety:
|
|
eps_eff = max(0.0, chall.epsilon - m)
|
|
cand = fgsm_untargeted(model, chall.x01, chall.label, eps_eff)
|
|
cand_pred = int(
|
|
torch.argmax(model(mnist_normalize(torch.from_numpy(cand))), dim=1).item()
|
|
)
|
|
if cand_pred != chall.label:
|
|
x_adv = cand
|
|
adv_pred = cand_pred
|
|
linf = float(np.max(np.abs(x_adv - chall.x01)))
|
|
break
|
|
|
|
if x_adv is None:
|
|
# Fall back to slight under-shoot of epsilon
|
|
eps_eff = chall.epsilon * 0.99
|
|
x_adv = fgsm_untargeted(model, chall.x01, chall.label, eps_eff)
|
|
adv_pred = int(
|
|
torch.argmax(model(mnist_normalize(torch.from_numpy(x_adv))), dim=1).item()
|
|
)
|
|
linf = float(np.max(np.abs(x_adv - chall.x01)))
|
|
|
|
print(
|
|
json.dumps({"linf": linf, "pred": adv_pred, "clean_pred": clean_pred}, indent=2)
|
|
)
|
|
|
|
# Submit
|
|
x2d = x_adv[0, 0]
|
|
b64 = png_from_x01(x2d)
|
|
r = requests.post(f"{args.host}/submit", json={"image_b64": b64}, timeout=10)
|
|
try:
|
|
r.raise_for_status()
|
|
except Exception:
|
|
print("Server response:", r.text)
|
|
raise
|
|
print("Flag:", r.json().get("flag"))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main() |