{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":2542588,"sourceType":"datasetVersion","datasetId":1541807}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader, random_split\nfrom torchvision import datasets, transforms, models\nfrom sklearn.metrics import accuracy_score, precision_recall_fscore_support\nimport numpy as np\nimport copy\nimport os\n\n# =============================================================================\n# 1. CONFIGURATION\n# =============================================================================\nclass Config:\n    # Auto-detect GPU\n    DEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n    # Federated Settings\n    NUM_CLIENTS = 5          # Number of virtual devices\n    GLOBAL_ROUNDS = 10       # Number of communication rounds\n    LOCAL_EPOCHS = 3         # Epochs per client per round\n    BATCH_SIZE = 32\n    LEARNING_RATE = 0.01\n    \n    # DATASET PATH (CHANGE THIS TO YOUR KAGGLE DATASET PATH)\n    # Example: '/kaggle/input/plantvillage-dataset/color'\n    DATA_DIR = '/kaggle/input/plantvillage-dataset/color' \n\n# =============================================================================\n# 2. DATA PREPARATION (The Fix: Split Train vs Test)\n# =============================================================================\ndef get_data_loaders(dataset_path, num_clients, batch_size):\n    \"\"\"\n    1. Loads the dataset.\n    2. Splits it 80% for Training (distributed to clients) and 20% for Testing (server only).\n    \"\"\"\n    # Standard ImageNet transforms\n    transform = transforms.Compose([\n        transforms.Resize((224, 224)),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n    ])\n    \n    # Handle missing dataset (Download CIFAR10 if PlantVillage path is wrong)\n    if os.path.exists(dataset_path):\n        print(f\"Loading data from {dataset_path}...\")\n        full_dataset = datasets.ImageFolder(root=dataset_path, transform=transform)\n    else:\n        print(f\"Path '{dataset_path}' not found. Downloading CIFAR10 for testing...\")\n        full_dataset = datasets.CIFAR10(root='./data', train=True, download=True, transform=transform)\n\n    # --- STEP 1: Split Full Data into TRAIN (80%) and TEST (20%) ---\n    train_size = int(0.8 * len(full_dataset))\n    test_size = len(full_dataset) - train_size\n    train_dataset, test_dataset = random_split(full_dataset, [train_size, test_size])\n    \n    # Create the Global Test Loader (Used for evaluation only)\n    test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)\n    \n    # --- STEP 2: Split TRAIN Data among Clients ---\n    # Divide the 80% training data into 'num_clients' parts\n    share_size = train_size // num_clients\n    remainder = train_size % num_clients\n    lengths = [share_size] * num_clients\n    lengths[-1] += remainder # Give remainder to last client\n    \n    client_datasets = random_split(train_dataset, lengths)\n    \n    client_loaders = []\n    for ds in client_datasets:\n        client_loaders.append(DataLoader(ds, batch_size=batch_size, shuffle=True))\n        \n    print(f\"Data Split: {train_size} Training images (distributed), {test_size} Test images (held-out).\")\n    return client_loaders, test_loader, len(full_dataset.classes)\n\n# =============================================================================\n# 3. METRICS (Precision, Recall, F1)\n# =============================================================================\ndef calculate_metrics(model, test_loader, device):\n    model.eval()\n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for images, labels in test_loader:\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            _, preds = torch.max(outputs, 1)\n            \n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n    \n    acc = accuracy_score(all_labels, all_preds)\n    precision, recall, f1, _ = precision_recall_fscore_support(\n        all_labels, all_preds, average='weighted', zero_division=0\n    )\n    return acc, precision, recall, f1\n\n# =============================================================================\n# 4. SERVER: Aggregation (FedAvg)\n# =============================================================================\ndef aggregate_weights(global_model, client_weights):\n    avg_weights = copy.deepcopy(global_model.state_dict())\n    for key in avg_weights.keys():\n        # Sum weights from all clients and divide by num_clients\n        key_sum = torch.stack([w[key] for w in client_weights], dim=0).sum(dim=0)\n        avg_weights[key] = key_sum / len(client_weights)\n    return avg_weights\n\n# =============================================================================\n# 5. CLIENT: Local Training\n# =============================================================================\ndef train_one_client(global_model, train_loader, epochs, lr, device):\n    local_model = copy.deepcopy(global_model)\n    local_model.to(device)\n    local_model.train()\n    \n    optimizer = optim.SGD(local_model.parameters(), lr=lr, momentum=0.9)\n    criterion = nn.CrossEntropyLoss()\n    \n    for epoch in range(epochs):\n        for images, labels in train_loader:\n            images, labels = images.to(device), labels.to(device)\n            \n            optimizer.zero_grad()\n            outputs = local_model(images)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n            \n    return local_model.state_dict()\n\n# =============================================================================\n# 6. MAIN SIMULATION LOOP\n# =============================================================================\ndef run_federated_learning():\n    print(f\"--- Starting Simulation on {Config.DEVICE} ---\")\n    \n    # 1. Setup Data\n    client_loaders, test_loader, num_classes = get_data_loaders(Config.DATA_DIR, Config.NUM_CLIENTS, Config.BATCH_SIZE)\n\n    # 2. Setup Model\n    # Note: Using weights=None to avoid Internet errors in Kaggle. \n    # If you have Internet enabled, you can change to weights='DEFAULT' for better results.\n    global_model = models.resnet18(weights=None) \n    global_model.fc = nn.Linear(global_model.fc.in_features, num_classes)\n    global_model.to(Config.DEVICE)\n    \n    # 3. Rounds Loop\n    for round_idx in range(Config.GLOBAL_ROUNDS):\n        print(f\"\\n=== Global Round {round_idx + 1}/{Config.GLOBAL_ROUNDS} ===\")\n        \n        local_weights = []\n        \n        # A. Train on Clients\n        for client_idx in range(Config.NUM_CLIENTS):\n            # print(f\"Training Client {client_idx + 1}...\", end='\\r')\n            w = train_one_client(global_model, client_loaders[client_idx], Config.LOCAL_EPOCHS, Config.LEARNING_RATE, Config.DEVICE)\n            local_weights.append(w)\n        \n        # B. Aggregate\n        print(f\"\\rAggregating weights from {Config.NUM_CLIENTS} clients...      \")\n        new_global_weights = aggregate_weights(global_model, local_weights)\n        global_model.load_state_dict(new_global_weights)\n        \n        # C. Evaluate on Test Set\n        acc, prec, rec, f1 = calculate_metrics(global_model, test_loader, Config.DEVICE)\n        \n        print(f\"Round {round_idx + 1} Test Metrics:\")\n        print(f\"  Accuracy:  {acc:.4f}\")\n        print(f\"  F1 Score:  {f1:.4f}\")\n        print(f\"  Precision: {prec:.4f}\")\n        print(f\"  Recall:    {rec:.4f}\")\n\nif __name__ == \"__main__\":\n    run_federated_learning()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-30T04:05:03.060138Z","iopub.execute_input":"2026-01-30T04:05:03.06041Z","iopub.status.idle":"2026-01-30T05:16:19.809709Z","shell.execute_reply.started":"2026-01-30T04:05:03.060385Z","shell.execute_reply":"2026-01-30T05:16:19.809085Z"}},"outputs":[],"execution_count":null}]}